torch.saveとtorch.loadの使い方

PyTorchにおけるモデルの保存と読み込みには主に torch.savetorch.load が用いられます。これらの関数は、モデルの学習済みパラメータや、モデル全体の状態をファイルとして保存・復元するために使用されます。


1. モデル保存の方法(torch.save

方法1:モデルのstate_dictのみ保存

最も一般的で推奨される方法です。

python
import torch # モデルの定義とインスタンス化 model = MyModel() # 学習済みパラメータの保存 torch.save(model.state_dict(), 'model_weights.pth')

ポイント:

  • state_dict() はモデルのパラメータ(重みやバイアスなど)を辞書型で返します。

  • 拡張子は慣習的に .pth.pt を使います。

  • モデルの構造(クラス定義など)は別途保持しておく必要があります。


方法2:モデル全体(構造+重み)を保存

python
torch.save(model, 'entire_model.pth')

注意点:

  • torch.save(model, ...) はPythonのpickleを使用してモデルのオブジェクト全体をシリアライズします。

  • モデルクラスの定義が保存時とまったく同じである必要があります。

  • 実務では再現性や互換性の面から、state_dict()方式が推奨されます。


2. モデル読み込みの方法(torch.load

方法1:state_dictから読み込む

python
# モデル定義 model = MyModel() # 学習済みパラメータの読み込み model.load_state_dict(torch.load('model_weights.pth')) # 評価モードに設定(推論時) model.eval()

補足:

  • model.eval() はドロップアウトやバッチ正規化などを評価モードに切り替えるための関数です。


方法2:モデル全体の読み込み

python
model = torch.load('entire_model.pth') model.eval()

注意点:

  • 保存時のモデルクラスの定義が必要です(インポート可能であること)。


3. 保存・読み込みにおける補足事項

  • GPUとCPUの切り替え:

    python
    # CPUで保存されたものをGPUに読み込む場合 model.load_state_dict(torch.load('model_weights.pth', map_location='cuda')) # GPUで保存されたものをCPUに読み込む場合 model.load_state_dict(torch.load('model_weights.pth', map_location='cpu'))
  • 複数の情報(モデル+オプティマイザなど)を保存することも可能:

    python
    torch.save({ 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'loss': loss, }, 'checkpoint.pth')

    読み込み時:

    python
    checkpoint = torch.load('checkpoint.pth') model.load_state_dict(checkpoint['model_state_dict']) optimizer.load_state_dict(checkpoint['optimizer_state_dict'])

まとめ

方法 保存関数 読み込み関数 特徴
推奨方法 model.state_dict()torch.save torch.loadload_state_dict 安全・柔軟・再現性が高い
簡易方法 torch.save(model) torch.load モデル構造も保存。構造の変更に弱い

PyTorchでは明示的にモデルの構造とパラメータを分離して扱うことが、再利用性やデバッグ性の観点から推奨されています。したがって、基本的には state_dict を保存・読み込みする方式を用いるべきです。

生成日:2025/05/22