GitHub Repo
PyTorchの事前コンパイル提案、カーネルをパッケージ化してリプレイ時のコンパイル禁止チェックを強化
著者によると、一連の変更により、推薦モデルの16個の分散プロセスがコンパイルを発生させずにリプレイを完了した。提案はまだドラフトで、他のモデルでの検証やデプロイ上の効果は確認されていない。

PyTorchの開発者は9月27日、リプレイに必要なカーネルもまとめてパッケージ化する事前コンパイルの提案を更新した。著者によると、一連の変更をFSDP2を使用する大規模な推薦モデルに適用したところ、16個の分散プロセスが100バッチを完了し、コンパイルは発生しなかった。確認時点で、PR #198789はまだドラフトであり、正式リリースされた機能ではない。[提案とテスト](https://github.com/pytorch/pytorch/pull/198789)
問題は、計算グラフを読み込んでも、基盤となるカーネルが準備済みとは限らないことだ。既存の公式ドキュメントでは、キャッシュは計算グラフ、Tritonのコンパイル結果、自動チューニングなどの層に分けられている。Mega-Cacheを使うと、これらの成果物をエクスポートして読み込めるが、PyTorchとTritonのバージョン、およびCUDA GPUが一致するかどうかがチェックされる。公式ドキュメントによると、自動チューニングでは候補となるカーネルを実測し、最速のものを選ぶ。再チューニングには依然としてコストがかかり、プロセス間でキャッシュを使い回す場合も、実行環境が制約を満たしているか確認する必要がある。[公式キャッシュドキュメント](https://docs.pytorch.org/tutorials/recipes/torch_compile_caching_tutorial.html)
新しい案では、3つのステップを提案している。`capture_runtime()`は実際に起動されたTritonカーネルとC++バイナリを記録し、`finalize_cache()`はカーネルランチャーとキャッシュの内容を確定する。`prepare_runtime()`はモデルの成果物を読み込む前に、カーネルを検証して準備する。Triton JITを直接使用する経路は、Tritonが実行時キャッシュのエクスポート用インターフェースを提供するかどうかにも左右される。[API設計](https://github.com/pytorch/pytorch/pull/198789)
併せて提案されているコンパイル禁止の仕組みでは、計算グラフのコンパイル、カーネルキャッシュミス、自動チューニングが発生すると、ただちにエラーを送出する。これにより、リプレイ中に気づかないままコンパイルが補われるのを防ぐ。エンジニアリングチームにとっては、「事前コンパイルが完了しているか」を確認可能な実行条件にできる。ただし、当初の関連PRはクローズされ、後続の提案に引き継がれており、インターフェースは今後変更される可能性がある。[関連設計](https://github.com/pytorch/pytorch/pull/198767)
デプロイの観点からは、選定したカーネルをモデルの成果物と一緒に管理することで、異なるプロセスが同じ計算経路を使っているか確認しやすくなり、準備漏れも早期に明らかになると考えられる。起動時の挙動を予測しやすくする必要があり、複数のワーカープロセスを同時に起動するデプロイでは、とりわけ関係の深いチェックだ。ただし、現時点の結果は著者による単一モデルのテストであり、すべてのモデルに当てはまる性能上の保証や数値的な保証ではない。今後は、上流でのレビュー、完全なCI、他のワークロードでの再現性、異なる入力形状をキャッシュでカバーできるかを確認する必要がある。正式にデプロイする前に、起動時間とキャッシュサイズも測定し、コンパイルコストの削減が成果物の転送や読み込みのオーバーヘッドで相殺されないか確かめる必要がある。[検証範囲](https://github.com/pytorch/pytorch/pull/198789)