ホームへ戻る

GitHub Repo

PyTorch開発版でクロージャの計算グラフの誤った再利用を報告、修正案で演算子チェックを追加

最小再現例では、加算と乗算のクロージャを順番にコンパイルすると、どちらも加算結果を返した。パッチは提案済みだが未マージで、公式CIでの検証、安定版への影響、事前コンパイルとの互換性は引き続き確認が必要だ。

Pytorch Deepdream (https://github.com/gordicaleksa/pytorch-deepdream) by gordicaleksa (https://github.com/gordicaleksa/pytorch-deepdream/commits?author=gordicaleksa) · MIT · Image source
zh-Hant

PyTorchコミュニティで9月20日、コンパイル結果の正確性に関する問題が報告された。同じファクトリ関数から生成された2つのクロージャが、テンソルの異なるメソッドをキャプチャしている場合、同じ計算グラフを誤って再利用する可能性がある。21日には修正案が提出されたが、確認時点ではまだマージされていない。この種の誤りは、一見正常な数値をそのまま返す可能性があり、動的な関数ファクトリを使うモデルのコードでは注意が必要だ。[問題報告](https://github.com/pytorch/pytorch/issues/197811)、[修正案](https://github.com/pytorch/pytorch/pull/197845)

最小再現例では、2つのクロージャがそれぞれテンソルの加算と乗算のディスクリプタをキャプチャする。6と3を入力すると、直接実行した場合は9と18を返すが、順番にコンパイルすると、どちらも9を返す。報告された環境は9月19日付の2.15開発版、CPU、Python 3.11で、報告者によると、独立した3回の実行すべてで再現し、警告も例外も発生しなかった。トップレベルで定義した加算・乗算関数に変更すると正常に動作した。[再現条件](https://github.com/pytorch/pytorch/issues/197811)

この問題は、コンパイルキャッシュの粒度に関係する。公式ドキュメントによると、`torch.compile`はコードオブジェクトごとに結果を保存するため、動的に作成された関数のコピーがキャッシュを共有する場合がある。再利用できるかどうかは、実行条件を確認するガードに依存する。クロージャがキャプチャした演算子がチェックされていなければ、テンソルの形状と型が一致していても、演算の意味が同じであるとは保証できない。[キャッシュの仕組み](https://docs.pytorch.org/docs/2.14/generated/torch.compile.html)

パッチ作者によると、従来の構築処理ではディスクリプタに対するガードが省略されていた。修正案では、メソッドディスクリプタとラッパーディスクリプタに同一性の照合を追加する。`BUILTIN_MATCH`を選んだ理由は、事前コンパイルに必要なガードのシリアライズにも対応するためだ。`ID_MATCH`だけでも同一性を照合できるが、保存の妨げになる。追加されたテストは、演算の切り替え、キャッシュの再利用、復元後に元の演算を受け入れ、別の演算を拒否することを確認する。[パッチの設計](https://github.com/pytorch/pytorch/pull/197845)

検証には依然として限界がある。作者はローカルでのテスト結果を提示しているが、自身の環境でクラッシュするアクセラレーターの同期テスト1件を除外している。公開ページでは公式CIの実行が承認待ちとなっており、レビュー結果もない。したがって、現段階では安定版が広く影響を受けるとも、正式な修正版がすでに存在するとも断言できない。[検証状況](https://github.com/pytorch/pytorch/pull/197845)

この事例を踏まえ、開発チームは同一プロセス内で異なる演算のクロージャを交互に実行し、呼び出しをまたいでキャッシュを保持したまま、未コンパイル時の結果と比較できる。これにより、呼び出し間のキャッシュ汚染を確認できる。特に、設定に基づいて複数のモデルバリエーションを生成し、常駐サービス内でコンパイルを繰り返すワークフローに適している。そのうえで、パッチのマージ、対応バージョン、事前コンパイル結果の復元テストを追跡し、アップグレードするか、独自に検証した代替の記述方法を採用するかを判断すべきだ。

出典

  1. Issue #197811: torch.compile silently reuses a graph when closures capture different Tensor method descriptors
  2. PR #197845: Guard on captured method/wrapper descriptors
  3. torch.compile — PyTorch 2.14 documentation