モデル保存の技術的課題と選定基準
時系列モデルのトレーニング完了後、以下のような問題に直面したことはありませんか?保存されたモデルファイルが巨大すぎてストレージコストが増加、特定のフレームワークに依存しすぎてデプロイが困難、異なるプラットフォーム間での互換性が欠如しているなど、モデル保存形式の選択は、開発効率から生産環境における安定性・パフォーマンスまでに影響を与えます。本記事では、5つの主要なモデル保存形式を比較し、時系列モデルライブラリ(TSLib)の実装に基づくベストプラクティスを紹介します。
主な保存形式とその特性
| 形式 | 保存内容 | ファイルサイズ | ロード速度 | フレームワーク互換性 | 主な用途 |
|---|---|---|---|---|---|
| .pth(state_dict) | パラメータのみ | 中程度 | 高速 | PyTorch専用 | 再トレーニング、微調整 |
| .pth(全体モデル) | アーキテクチャ+パラメータ | やや大きい | 中程度 | PyTorch専用 | プロトタイプ開発 |
| .pt(TorchScript) | 静的グラフ+パラメータ | 中程度 | 非常に高速 | C++、PyTorch | 組み込み、モバイル |
| .onnx | 標準化されたモデル構造 | 中程度 | 中程度 | 多フレームワーク対応 | クロスプラットフォーム |
| .h5(HDF5) | メタデータ含む | やや大きい | 遅め | Keras/TensorFlow | データ交換 |
各形式の実装例と比較
state_dict形式(推奨)
モデルのパラメータのみを保存する形式。TSLibでは以下のようなコードで実装されています:
class EarlyStopping:
def save_checkpoint(self, val_loss, model, path):
torch.save(model.state_dict(), path + '/' + 'checkpoint.pth')
self.val_loss_min = val_loss
保存とロードの例:
# 保存
torch.save(model.state_dict(), 'timesnet_checkpoint.pth')
# ロード
model = TimesNet(args)
model.load_state_dict(torch.load('timesnet_checkpoint.pth'))
model.eval()
ONNX形式でのエクスポート
他フレームワークへの移行を容易にするための形式:
dummy_input = (
torch.randn(1, 96, 7).to(device),
torch.randn(1, 96, 4).to(device),
torch.randn(1, 192, 7).to(device),
torch.randn(1, 192, 4).to(device)
)
torch.onnx.export(
model,
dummy_input,
"timesnet.onnx",
input_names=["batch_x", "batch_x_mark", "dec_inp", "batch_y_mark"],
output_names=["output"],
dynamic_axes={
"batch_x": {0: "batch_size"},
"output": {0: "batch_size"}
}
)
TorchScript形式
C++や組み込み環境での高速推論に適した形式:
scripted_model = torch.jit.script(model)
torch.jit.save(scripted_model, "timesnet_scripted.pt")
# 推論
loaded_model = torch.jit.load("timesnet_scripted.pt")
output = loaded_model(batch_x, batch_x_mark, dec_inp, batch_y_mark)
拡張保存戦略の提案
単なるパラメータ保存に加え、設定情報や前処理統計情報を含む「モデルパッケージ」形式が推奨されます:
def save_model_complete(model, args, stats, performance, path):
os.makedirs(path, exist_ok=True)
torch.save(model.state_dict(), os.path.join(path, "model.pth"))
with open(os.path.join(path, "config.json"), "w") as f:
json.dump(vars(args), f, indent=4)
with open(os.path.join(path, "data_stats.pkl"), "wb") as f:
pickle.dump(stats, f)
with open(os.path.join(path, "performance.txt"), "w") as f:
f.write(f"MAE: {performance['mae']:.4f}\n")
f.write(f"MSE: {performance['mse']:.4f}\n")
f.write(f"Training Time: {performance['train_time']:.2f}s\n")
今後の改善と代替手法
- モデル量子化:サイズ削減と速度向上を実現
- 分散保存:マルチGPU環境向けの保存戦略
- 増分保存:定期的なチェックポイント管理