nn.Moduleを継承したカスタムモデルの作成

PyTorchにおけるモデル構築の基本は、torch.nn.Moduleを継承してカスタムモデルを定義することです。この方法により、ニューラルネットワークの構造や前向き伝播(forward)処理を柔軟に設計できます。以下では、「nn.Moduleを継承したカスタムモデルの作成」について、詳しく説明します。


1. nn.Moduleの役割

torch.nn.Moduleは、PyTorchのすべてのニューラルネットワークモデルの基本クラスです。これを継承することで、以下の機能が提供されます:

  • パラメータ(nn.Linearなど)の自動登録と管理

  • .to(device).cuda().eval().train() などの便利なメソッド

  • モデルの保存・読み込みが容易になる(state_dict()の利用)


2. 基本的な構成

カスタムモデルを作成するには、次の手順を踏みます:

a. __init__() メソッドで層を定義

b. forward() メソッドでデータの流れ(前向き伝播)を定義


3. 実装例(全結合2層のMLP)

python
import torch import torch.nn as nn import torch.nn.functional as F class MyModel(nn.Module): def __init__(self, input_size, hidden_size, output_size): super(MyModel, self).__init__() self.fc1 = nn.Linear(input_size, hidden_size) # 入力層→隠れ層 self.fc2 = nn.Linear(hidden_size, output_size) # 隠れ層→出力層 def forward(self, x): x = F.relu(self.fc1(x)) # 活性化関数ReLU x = self.fc2(x) return x

4. モデルの使用方法

python
# インスタンス化 model = MyModel(input_size=100, hidden_size=50, output_size=10) # 入力データ x = torch.randn(16, 100) # バッチサイズ16、特徴量100 # 推論 output = model(x)

5. モデルの拡張性

nn.Moduleを継承することで、以下のような高度な構成も簡単に実現できます:

  • 畳み込みニューラルネットワーク(CNN)

  • 再帰型ニューラルネットワーク(RNN/LSTM)

  • モジュールの再利用(nn.Sequentialself.block = nn.ModuleList([...]) など)

  • 条件分岐・ループによる柔軟なforward処理


6. モデルの保存と読み込み

python
# 保存 torch.save(model.state_dict(), 'model.pth') # 読み込み model.load_state_dict(torch.load('model.pth')) model.eval()

まとめ

nn.Moduleを継承することで、モデルの構造や動作を自由に定義でき、PyTorchのエコシステムとの連携もスムーズに行えます。特に複雑なモデル設計や研究開発においては、この方法が推奨されます。

生成日:2025/05/22