回帰モデル

PyTorchにおける回帰モデル(regression model)は、入力データから連続値(実数値)を予測するためのモデルです。たとえば、家の価格、温度、株価などを予測するタスクで利用されます。分類モデルとは異なり、出力はカテゴリではなくスカラーやベクトルの実数値となります。

以下では、PyTorchで回帰モデルを構築・学習・評価するための主要なステップを詳しく解説します。


1. データの準備

回帰モデルでは、入力 X に対して実数値のターゲット y をペアで持つデータセットを用意します。

python
import torch from torch.utils.data import Dataset class RegressionDataset(Dataset): def __init__(self, X, y): self.X = torch.tensor(X, dtype=torch.float32) self.y = torch.tensor(y, dtype=torch.float32) def __len__(self): return len(self.X) def __getitem__(self, idx): return self.X[idx], self.y[idx]

2. モデルの定義(nn.Module)

回帰モデルでは、通常、最後の層に活性化関数を使わずに出力することで、制約のない連続値を予測します。

python
import torch.nn as nn class RegressionModel(nn.Module): def __init__(self, input_dim): super().__init__() self.linear = nn.Sequential( nn.Linear(input_dim, 64), nn.ReLU(), nn.Linear(64, 1) # 出力は1次元(スカラー) ) def forward(self, x): return self.linear(x)

3. 損失関数とオプティマイザの設定

回帰問題では主に**平均二乗誤差(MSE)**が用いられます。

python
model = RegressionModel(input_dim=10) criterion = nn.MSELoss() optimizer = torch.optim.Adam(model.parameters(), lr=0.001)

4. 学習ループ

エポック毎に予測を行い、損失を計算してパラメータを更新します。

python
for epoch in range(100): model.train() total_loss = 0.0 for X_batch, y_batch in dataloader: optimizer.zero_grad() outputs = model(X_batch).squeeze() # 出力の形状調整 loss = criterion(outputs, y_batch) loss.backward() optimizer.step() total_loss += loss.item() print(f"Epoch {epoch}: Loss = {total_loss:.4f}")

5. 評価指標

回帰モデルの評価には以下がよく使われます:

  • MSE(Mean Squared Error)

  • RMSE(Root Mean Squared Error)

  • MAE(Mean Absolute Error)

  • R²スコア(決定係数)

python
from sklearn.metrics import mean_squared_error, r2_score model.eval() with torch.no_grad(): y_pred = model(X_test_tensor).squeeze().numpy() y_true = y_test_tensor.numpy() print("MSE:", mean_squared_error(y_true, y_pred)) print("R²:", r2_score(y_true, y_pred))

6. 応用例

  • 不動産価格の予測(入力: 間取り・築年数・面積 → 出力: 価格)

  • 気温予測(入力: 時刻・場所・過去の温度 → 出力: 将来の温度)

  • 売上予測(入力: プロモーションデータ・季節性 → 出力: 数値)


補足:出力層と活性化関数の注意点

分類モデルと異なり、回帰モデルでは出力層に活性化関数を使わないのが基本です。ただし、特定の範囲に出力を制限したい場合は、sigmoidsoftplus などを使うこともあります。

生成日:2025/05/22