画像分類(CIFAR-10、MNISTなど)

MXNetにおける画像分類(CIFAR-10、MNISTなど)の実装は、Gluon APIを用いることで柔軟かつ効率的に行えます。以下では、典型的な画像分類タスクの流れに沿って、主要なステップを詳しく説明します。


1. データの読み込みと前処理

MXNetでは、gluon.data.visionモジュールにMNISTやCIFAR-10といった一般的なデータセットが組み込まれています。

python
from mxnet.gluon.data.vision import transforms from mxnet.gluon.data.vision import MNIST transformer = transforms.Compose([ transforms.ToTensor(), # [0, 255] を [0, 1] に正規化 transforms.Normalize(0.13, 0.31) # 平均と標準偏差で標準化(MNIST用) ]) train_data = MNIST(train=True).transform_first(transformer) test_data = MNIST(train=False).transform_first(transformer)

DataLoaderを使ってバッチ化します。

python
from mxnet.gluon.data import DataLoader batch_size = 64 train_loader = DataLoader(train_data, batch_size=batch_size, shuffle=True) test_loader = DataLoader(test_data, batch_size=batch_size, shuffle=False)

2. モデルの定義

gluon.nn.Sequentialを使って畳み込みニューラルネットワーク(CNN)を定義します。

python
from mxnet.gluon import nn net = nn.Sequential() net.add( nn.Conv2D(channels=32, kernel_size=3, activation='relu'), nn.MaxPool2D(pool_size=2), nn.Conv2D(channels=64, kernel_size=3, activation='relu'), nn.MaxPool2D(pool_size=2), nn.Flatten(), nn.Dense(128, activation='relu'), nn.Dense(10) # 出力クラス数 ) net.initialize()

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

python
from mxnet import gluon loss_fn = gluon.loss.SoftmaxCrossEntropyLoss() trainer = gluon.Trainer(net.collect_params(), 'adam', {'learning_rate': 0.001})

4. 学習ループの実装

python
from mxnet import autograd, nd def train_model(epochs): for epoch in range(epochs): total_loss = 0 for data, label in train_loader: with autograd.record(): output = net(data) loss = loss_fn(output, label) loss.backward() trainer.step(batch_size) total_loss += loss.mean().asscalar() print(f'Epoch {epoch + 1}, Loss {total_loss / len(train_loader):.4f}')

5. モデルの評価

python
def evaluate_accuracy(data_iter): acc = mx.metric.Accuracy() for data, label in data_iter: output = net(data) predictions = nd.argmax(output, axis=1) acc.update(label, predictions) return acc.get()[1] test_acc = evaluate_accuracy(test_loader) print(f'Test Accuracy: {test_acc:.4f}')

6. CIFAR-10への応用

MNISTと同様の手順で、以下の点を変更します:

  • データセットを CIFAR10 に変更

  • 入力チャネルが3(RGB)のため、モデルの入力に合わせてネットワーク構造を調整

  • 標準化の平均・標準偏差もCIFAR-10用に変更(例:transforms.Normalize([0.4914, 0.4822, 0.4465], [0.2470, 0.2435, 0.2616])


補足:GPUの利用

GPUがある場合は、以下のようにしてコンテキストを指定します。

python
from mxnet import gpu ctx = gpu() # GPUがなければ mx.cpu() net.initialize(ctx=ctx)

学習・推論のときはデータとモデルを ctx に送る必要があります。


まとめ

MXNetにおける画像分類の実装は以下の構成を通じて行います:

  1. データセットの読み込みと前処理

  2. CNNモデルの構築

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

  4. 学習ループの実行

  5. テスト精度の評価

これらの手順を組み合わせることで、MNISTやCIFAR-10といった代表的なデータセットに対する画像分類タスクを効率的に実装できます。さらに、HybridBlockの利用による高速化やモデルのエクスポートにも発展できます。

生成日:2025/05/23