0Pricing
Learn AI with Python · レッスン

カスタムデータセットとDataLoader

torch.utils.data.Dataset、__len__/__getitem__、DataLoader、変換、データ拡張について学習します。

「カスタムデータセットとDataLoader」はCoddyKit上の無料Learn AI with Pythonレッスンです。 これはレッスン2/4です。 下記で完全なレッスンを無料で読むことができます。その後、ブラウザ内の組み込みコードエディタと24時間対応のAIチューターでハンズオン演習できます。 これはLearn AI with Python学習パスの一部であり、ウェブとCoddyKitアプリ全体で進捗が同期されます。 Learn AI with Pythonコースには全4レッスンが含まれています。

モデルへのデータ入力

トレーニングには、データを効率的に読み込み、変換し、バッチ化する方法が必要です。PyTorchには二つの抽象化機能があります。Datasetは1つのサンプルを取得する方法を把握し、DataLoaderはサンプルをバッチ化してシャッフルします。

from torch.utils.data import Dataset, DataLoader

Datasetインターフェース

カスタムDatasetサブクラスでは、二つのメソッドを実装する必要があります。__len__はサンプル数を返し、__getitem__はインデックスに対応するサンプルを返します。PyTorchはこれらを呼び出してデータを取得します。

__len__の実装

__len__はデータセットのサイズをPyTorchに伝えます。これにより、存在するインデックスの数や、1エポックに含まれるバッチ数が決まります。

class ImageDataset(Dataset):
    def __init__(self, paths, labels):
        self.paths = paths
        self.labels = labels

    def __len__(self):
        return len(self.paths)

__getitem__の実装

__getitem__は、インデックスを受け取って1つのサンプル(およびそのラベル)を読み込み、返します。画像ファイルを開き、テンソルに変換する処理はここに記述します。

from PIL import Image

    def __getitem__(self, idx):
        img = Image.open(self.paths[idx]).convert("RGB")
        label = self.labels[idx]
        return img, label

なぜTransformsを使うのか

生の画像はサイズや画素値の範囲がそれぞれ異なります。Transformsを使うと、固定サイズへのリサイズ、テンソルへの変換、画素値の正規化を行ってデータを標準化できるため、モデルを安定してトレーニングできます。

from torchvision import transforms

transforms.Compose

transforms.Composeは複数の変換を1つのパイプラインに連結し、順番に適用します。一般的な連鎖は、Resize、ToTensor、Normalizeの順です。

tf = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406],
                         std=[0.229, 0.224, 0.225])
])

Transformsの仕組み

Resizeは空間サイズを固定し、ToTensorはPIL画像をテンソルに変換して画素値を[0,1]にスケーリングします。Normalizeは各チャンネルを平均0、分散1になるように移動・スケーリングし、収束を速めます。

DatasetでTransformsを適用する

Transformをデータセットに渡し、__getitem__の中で適用します。これにより、各サンプルを取得するたびに一貫した前処理を行えます。

class ImageDataset(Dataset):
    def __init__(self, paths, labels, transform):
        self.paths, self.labels, self.transform = paths, labels, transform

    def __getitem__(self, idx):
        img = Image.open(self.paths[idx]).convert("RGB")
        return self.transform(img), self.labels[idx]

DataLoaderでラップする

DataLoaderはDatasetをバッチの反復可能なオブジェクトに変換します。batch_sizeで1ステップあたりのサンプル数を設定し、shuffle=Trueで各エポックの順序をランダム化します(トレーニングでは重要です)。

dataset = ImageDataset(paths, labels, tf)
loader = DataLoader(dataset, batch_size=32, shuffle=True)

高速化のためのnum_workers

num_workersは、GPUがトレーニングしている間にデータの読み込みと変換を行う並列サブプロセスを起動し、I/Oの待ち時間を隠します。4などの値を設定すると、GPUを待機させずにデータを供給できることがよくあります。

loader = DataLoader(
    dataset,
    batch_size=32,
    shuffle=True,
    num_workers=4
)

バッチの反復処理

DataLoaderをループして、バッチ化されたテンソルを取得します。各反復では(images, labels)が返されます。imagesの形状は[batch_size, channels, H, W]で、モデルにそのまま入力できます。

for images, labels in loader:
    print(images.shape)  # torch.Size([32, 3, 224, 224])
    break

理解度チェック

データパイプラインの理解度を確認しましょう。

復習:DatasetsとDataLoaders

__len__と__getitem__を備えたカスタムDatasetを作成し、transforms.Compose(Resize、ToTensor、Normalize)で画像を前処理しました。さらに、batch_size、shuffle、num_workersを設定したDataLoaderでラップし、モデルにバッチを効率的に供給しました。

よくある質問

「カスタムデータセットとDataLoader」レッスンは無料ですか?

はい。「カスタムデータセットとDataLoader」の完全なテキストはこのウェブで無料で読めます。インタラクティブに演習し(組み込みコードエディタと24時間対応のAIチューター)、Learn AI with Pythonコースの残りをアンロックするには、CoddyKit PROにアップグレードしてください。 Learn AI with Pythonコースには全4レッスンが含まれています。

「カスタムデータセットとDataLoader」で何を学びますか?

torch.utils.data.Dataset、__len__/__getitem__、DataLoader、変換、データ拡張について学習します。 ブラウザで直接実行するハンズオンコードでLearn AI with Pythonを演習し、24時間対応のAIチューターがレッスンを進める中での質問に答えます。

Learn AI with Pythonを始めるのに経験は必要ですか?

事前経験は必要ありません。CoddyKitのLearn AI with Pythonは初級者から上級者向けに構成されているため、ここから始めるか最初から始めて、自分のペースで進むことができます。 これはレッスン2/4です。

「カスタムデータセットとDataLoader」レッスンにはどのくらい時間がかかりますか?

ほとんどのCoddyKitレッスンは約5~10分かかります。各レッスンはコンパクトでインタラクティブなので、着実に進歩し、ウェブとアプリ全体で正確に前回の場所から再開できます。

このLearn AI with Pythonレッスンでコードを書いて実行できますか?

はい。すべてのLearn AI with Pythonレッスンに組み込みコードエディタが含まれているため、ブラウザでリアルコードを書いて実行し、即座のAIフィードバックを取得できます。ローカル設定は不要です。

このコースのすべてのレッスン

  1. PyTorchのテンソルとAutograd
  2. カスタムデータセットとDataLoader
  3. PyTorchでCNNを構築・学習する
  4. YOLOv8による物体検出
← Learn AI with Pythonに戻る