分類モデル(MNIST、CIFAR-10など)

PyTorchにおける分類モデル(classification model)の構築は、画像やテキストなどのデータを複数のカテゴリに分類するタスクに用いられます。ここでは代表的な画像分類タスクであるMNIST(手書き数字の分類)やCIFAR-10(日常物体の分類)を例に、実践的な分類モデルの構築プロセスを解説します。


1. 問題設定

  • MNIST: 28×28のグレースケール画像で、0〜9の手書き数字に分類。

  • CIFAR-10: 32×32のカラー画像で、10クラス(飛行機、自動車、鳥、猫など)に分類。


2. 全体フロー

  1. データの準備(torchvision.datasetsDataLoader

  2. モデルの構築(torch.nn.Module

  3. 損失関数と最適化手法の定義

  4. 学習ループの実装

  5. 評価・検証


3. データの準備

python
import torchvision.transforms as transforms import torchvision.datasets as datasets from torch.utils.data import DataLoader # MNIST用(グレースケール) transform_mnist = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) # CIFAR-10用(カラー) transform_cifar10 = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform_mnist) test_dataset = datasets.MNIST(root='./data', train=False, download=True, transform=transform_mnist) train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True) test_loader = DataLoader(test_dataset, batch_size=64, shuffle=False)

4. モデルの構築

MNIST用のシンプルなCNN(畳み込みニューラルネットワーク)

python
import torch.nn as nn import torch.nn.functional as F class MNISTClassifier(nn.Module): def __init__(self): super(MNISTClassifier, self).__init__() self.conv1 = nn.Conv2d(1, 32, kernel_size=3) # 入力チャネル1(グレースケール) self.conv2 = nn.Conv2d(32, 64, kernel_size=3) self.fc1 = nn.Linear(9216, 128) self.fc2 = nn.Linear(128, 10) # 出力クラス数10 def forward(self, x): x = F.relu(self.conv1(x)) x = F.relu(self.conv2(x)) x = F.max_pool2d(x, 2) x = x.view(-1, 9216) # Flatten x = F.relu(self.fc1(x)) x = self.fc2(x) return x

5. 損失関数と最適化手法

python
import torch.optim as optim model = MNISTClassifier() criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=0.001)

6. 学習ループの実装

python
for epoch in range(10): model.train() for images, labels in train_loader: optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() print(f"Epoch {epoch+1} completed.")

7. 評価・検証

python
correct = 0 total = 0 model.eval() with torch.no_grad(): for images, labels in test_loader: outputs = model(images) _, predicted = torch.max(outputs.data, 1) total += labels.size(0) correct += (predicted == labels).sum().item() print(f"Accuracy: {100 * correct / total:.2f}%")

8. CIFAR-10の場合の補足

  • CIFAR-10はカラー画像なので、Conv2dの入力チャネルは3になります。

  • 入力サイズが異なるため、Linear層のユニット数も変更する必要があります。

  • データ拡張(transforms.RandomHorizontalFlipなど)を加えることで、汎化性能の向上が期待できます。


まとめ

分類モデルの構築では以下の点が重要です:

  • 入力データの前処理(正規化、サイズ調整など)

  • モデル構造の設計(畳み込み、全結合層)

  • 適切な損失関数(多クラス分類ならCrossEntropyLoss)

  • オプティマイザ(SGD, Adamなど)の選択

  • 学習ループと評価指標の設計

PyTorchはこれらすべてを柔軟かつ明快に記述できるため、実験的・実践的なモデル構築に非常に適しています。

生成日:2025/05/22