SparseDecoding、生成段階に適応したLLM枝刈りを提案
導入
大規模言語モデルの推論では、計算量だけが遅延を決めるわけではない。デコード段階では通常、モデルは一度に1トークンを生成する一方、そのたびに大量の重みを読み出す。このため、メモリ帯域とデータ移動が大きなボトルネックになる。枝刈りによって非ゼロパラメータを減らせば読み出し量を抑えられるが、実際の速度向上には、枝刈りの基準と実行カーネルがデコード処理に合っていることも必要になる。
Westlake ENCODE Labの研究チームは、この課題に対して SparseDecoding を提案した。アルゴリズム面では、固定された自然文ではなく、密なモデルが自己回帰的に生成する過程の活性値を使って枝刈りを校正する。システム面では、トークン単位の推論で重要になる疎行列ベクトル積(SpMV)を対象にしたN:M疎カーネルを開発した。Llama-3.1-8B、Llama-3.3-70B、Qwen3-14B、Qwen3-32Bを含む評価で、A100 GPU上で最大1.48倍のエンドツーエンドデコード高速化を報告している。
従来の校正が抱える分布のずれ
学習を追加で行わない枝刈り手法の中には、あらかじめ用意した自然文からHessian関連の情報を計算し、重みの重要度を推定するものがある。この方法は扱いやすいが、自然文で得られた活性値が生成時の活性値を十分に代表するという前提に立っている。
しかし、自己回帰デコードでモデルに入力されるのは、人間が固定した系列だけではなく、モデル自身が前のステップで生成したトークンである。生成の進行に伴って隠れ状態や層ごとの活性分布が変化し、固定テキストによる校正データから離れる可能性がある。研究では、この差が枝刈りの目的と実際の生成過程の不一致を招き、枝刈り後のモデル性能を損なう要因になると説明している。
SparseDecodingは、密なモデルが自己回帰生成を行う間に各層の活性値を収集する。prefill段階を除外することで、対象となる反復的なトークン生成の分布に校正を寄せる。つまり、どの重みが一般的な文章に重要かだけでなく、モデルが連続してトークンを生成する際に重要かを評価する設計である。
主なポイント
- デコード対応の校正:固定自然文ではなく、密なモデルの生成中に得られる層別活性値を利用する。
- SpMVを重視:デコードで支配的な疎行列ベクトル積を対象にし、疎行列行列積(SpMM)中心の既存最適化を補う。
- アルゴリズムとシステムの協調:ビットマスク索引と固定ステップ走査を備えたN:M疎SpMVカーネルを実装する。
- 複数モデルで評価:LlamaおよびQwenの複数規模モデルを、長文生成ベンチマークで比較した。
意義と注意点
パラメータ数や非ゼロ要素を減らすだけでは、実測のレイテンシ低下は保証されない。ランタイムがベクトル単位の疎なアクセスを効率よく処理できなければ、理論上の削減量が実際の速度に反映されにくいからだ。SparseDecodingは、校正データの分布とSpMV実装を同じデプロイ課題として扱っている点に意義がある。
一方、報告された速度向上は、特定のモデル、N:M疎性、A100 GPUに基づく。実際の効果は、ハードウェア、バッチサイズ、系列長、推論フレームワークの実装によって変わり得る。それでも、主な用途が逐次的なトークン生成であるシステムにとって、生成そのものから校正信号を取るという発想は、単純に疎性を高める方法とは異なる実用的な方向性を示している。
コメント
ログイン状態を確認中…
コメントを読み込み中…