時系列モデル保存形式の比較とベストプラクティス

モデル保存の技術的課題と選定基準

時系列モデルのトレーニング完了後、以下のような問題に直面したことはありませんか?保存されたモデルファイルが巨大すぎてストレージコストが増加、特定のフレームワークに依存しすぎてデプロイが困難、異なるプラットフォーム間での互換性が欠如しているなど、モデル保存形式の選択は、開発効率から生産環境における安定性・パフォーマンスまでに影響を与えます。本記事では、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環境向けの保存戦略
  • 増分保存:定期的なチェックポイント管理

タグ: PyTorch ONNX TorchScript 時系列モデル モデル保存

8月18日 15:22 投稿