GoogleのGemma-4 31BモデルをAWS Inferentia2 inf2.24xlargeに移植したところ、CPUリファレンスとトークン単位で完全に一致した。しかし、生成された文章はすべて意味不明なものだった。「一致していること」と「動作していること」の間のこの乖離は、巨大なLLMをAmazonのカスタム推論チップに詰め込もうとするすべての人への警告となる。

なぜトークン単位の一致だけでは不十分なのか

開発者は、Inferentiaデバイスからの各出力トークンを、モデルのCPU実行によって生成されたトークンと比較した。ストリームは同一であったため、ハードウェアはリファレンス実装を正確に再現しているように見えた。しかし実際には、両方のストリームは、チャットテンプレートが削除され、誤ったターンマーカーが供給された不完全なプロンプトをモデルに送り込んでいた。テンプレートが欠落していたため、モデルは無限ループに陥り、支離滅裂な内容を吐き出した。ハードウェアは、リファレンスコードに存在していたバグを正確に再現するという、その役割を果たしてしまったのだ。

教訓は単純だ。SEQ_MATCH(逐次的なトークンの等価性)は、正当性と等価ではない。もしリファレンス実装が壊れていれば、忠実なハードウェアの複製も同じ失敗を引き継いでしまう。検証はトークンレベルのパリティ(一致)を超え、適切にフォーマットされた入力を用いたエンドツーエンドの機能チェックを行う必要がある。

パラメータを装うバッファ

ロードフェーズ中に、モデルローダーはlayer_scalarと呼ばれるコンポーネントをスキップした。コード上、このオブジェクトはPyTorchのモデル定義においてparameter(パラメータ)ではなくbuffer(バッファ)として登録されていた。バッファとは、学習によって更新されない静的なテンソルであり、多くのローダーはNeuron互換形式に変換する際にこれらを無視してしまう。これをスキップしたことで、いくつかのレイヤーのスケーリング係数がデフォルト値のままとなり、ネットワーク全体の計算が歪んでしまった。エラーは発生せず、モデルのコンパイルも推論パイプラインの実行も行われたが、数値結果が狂っていたのである。

大規模モデルをInferentiaに移行する際は、パラメータ以外のすべてのテンソルを監査すべきだ。たとえ学習対象ではないテンソルであっても、正しいフォワードパスの計算に不可欠な場合がある。バッファが含まれているかを手動で検証することで、診断が困難なサイレントなスケールエラーを防ぐことができる。

スポットインスタンスの揮発性と39分間のコンパイル

310億パラメータのモデルをスポットインスタンスで実行するのは安価に見えるが、その節約には予測不可能なインスタンスの回収(reclaim)というリスクが伴う。モデルをNeuron互換コードに変換するためのコンパイル時間(約39分間)は、AWSがインスタンスを回収した瞬間に消失した。中断に耐えるため、開発者は3層のセーフティネットを構築した。

  • ModelBuilder:メモリ使用量をホストの制限である384 GB以内に抑え、再起動を余儀なくされるクラッシュを回避した。
  • S3への即時ミラーリング:生の重みファイルとコンパイル済みの「neffs」(Neuron実行ファイル)の両方をS3にミラーリングすることで、新しいインスタンスが前回の続きから正確に再開できるようにした。
  • マルチリージョン・ポリャー(poller):AWSリージョンをスキャンして利用可能なスポット容量を探し、空きが出次第すぐに新しいインスタンスを起動した。

これらのステップにより、脆弱で単一障害点となっていたコンパイル作業が、スポット市場の変動に耐えうる弾力的なパイプラインへと変わった。

混合アテンションレイアウトによるシャーディングの落とし穴

Gemma-4 31Bは、2種類のアテンション構成を使用している。一部のレイヤーは4つのKV(Key-Value)ヘッドを採用しているが、他のレイヤーは異なる数を使用している。レイヤーのKVヘッド数がきれいに割り切れない場合、モデルを8つの並列ランクに均等に分割(シャーディング)することはできない。4ヘッドのレイヤーを8つのランクにシャーディングしようとすると、各ランクが「0.5ヘッド」を処理しなければならなくなり、これは数学的に不可能であり、シェイプの不一致やランタイムエラーを引き起こす。

解決策は、グローバルにシャーディングされたレイヤー(ヘッド数が互換性のあるもの)をすべてのランクに複製し、ヘッド数が均等に分割できる「スライディング」レイヤーのみをシャーディングすることだった。このハイブリッド戦略により、KVヘッドの不正な分割を避けつつ、テンソル並列の効率を維持し、以前の試行を悩ませていたテンソル並列化エラーを排除することができた。

まとめ

巨大なLLMをInferentiaに移植することは、単なる「コンパイルして実行する」だけの作業ではない。トークンの一致を超えた厳格な機能テスト、パラメータかバッファかを問わずすべてのテンソルが正しく処理されているかの細かな検証、そしてスポットインスタンスの回収を想定したデプロイ戦略が求められる。最後に、シャーディングはモデル内部のアテンション幾何学(geometry)を尊重しなければならない。さもなければ、スピードを約束するはずの並列化が、サイレントな失敗の源となってしまう。