この記事は軽量化シリーズの第4回(#4)です。量子化 #2で「8ビットはタダ、2ビットは崩壊」まで見たなら、今回はなぜ崩壊するのか、そしてどう蘇らせるのかを見ます。コードは github.com/warpspaceinc/efficient-ml-practicelow-bit-quantization.ipynb にあり、以下のすべてのplotと表はそのコードを実際に回して測定した我々の数値です。

FP4という無理筋

NVIDIA Blackwellのスペック表には目を疑う行があります。FP4 tensor coreがFP8のちょうど2倍のスループットを出すというのです。ビットを半分にするたびにスループットが2倍になる構造なので当然の算数なのですが、問題はFP4(E2M1)というフォーマットそのものです。符号1ビット + 指数2ビット + 仮数1ビット。表現可能な値はたった15個です。

$$\pm\{0.5,\ 1,\ 1.5,\ 2,\ 3,\ 4,\ 6\}\ \cup\ \{0\}$$

FP4 vs INT4の格子 — FP4は0付近が密で端に行くほど疎

INT4の均等格子と違い、FP4は0付近が密です(間隔0.5、端は2.0)。重みが0付近に釣鐘型に集まることを考えれば合理的な配置です。しかし15個は15個です。

表現範囲が[−6, 6]しかない点も重要です。実数の重みをこの狭い窓に押し込むには、scale factor $S$ で割って縮める必要があります。そしてこのscaleをどう決めるかが低ビット量子化の生死を分けます。この記事の残りはすべてその話です。

実験は#2と同じMNIST MLP(784→256→128→10、FP32精度97.4%)で行います。


1. 犯人:アウトライア1個が層を消す

最も単純なscale設計は、テンソル全体に1個だけ使うper-tensorです。絶対値が最大の重みがFP4の最大値6に来るように決めます:$S = |W|_{\max} / 6$。何も切り捨てられないので安全に見えます。

しかしこの式をよく見ると、恐ろしい点が1つあります。格子全体の間隔を、最大値たった1つが決めているのです。最大値が大きくなればscaleも大きくなり、格子が疎になって、0付近に集まる大多数の重みの解像度が悪くなります。では、何らかの理由で異常に大きい値、つまりアウトライア(outlier)が1つでもあったら?

確かめてみましょう。fc1の重み(20万個)のコピーで要素をたった1個だけ最大値の50倍に膨らませ、FP4量子化後に残りの重みが受ける被害を測りました。

アウトライア1個注入時のgranularity別MSE(logスケール)

scaleの単位MSE(正常)MSE(アウトライア注入)増加率
per-tensor5.4e-051.2e-0322倍
per-channel1.7e-052.3e-051倍
per-group(32)1.2e-051.2e-051倍

per-tensorの「22倍」は数字以上に中身が悲惨です。順を追うとこうなります。アウトライアが $|W|_{\max}$ を50倍に引き上げるとscaleも50倍になります。すると0でない最小のFP4レベル($0.5 \times S$)が、正常な重みすべて($|w| \le 0.3$)より大きくなります。正常な重みから見れば、最も近い格子点は0になってしまう。つまり20万個の重みがほぼすべて0にスナップされます。アウトライア1個のせいで層が丸ごと消えたのです。

これは人工的なシナリオだけの話ではありません。MIT 6.5940の講義1が示す実例はMobileNetV2で、最初のdepthwise層はチャネル間の重み範囲が100倍以上違います。低ビット量子化の失敗の大半はこのパターンです。アウトライア1個がscaleを汚染し、残り全体の解像度を殺すのです。

2. 処方 ①:scaleを細かく刻む — グループ量子化

上の表がすでに解決策も示しています。問題の本質が「アウトライア1個がテンソル全体のscaleを汚染する」ことなら、scaleが担当する範囲を狭めればいい。被害の爆発半径を縮めるのです。

  • per-channel:チャネル(行)ごとにscaleを1個。アウトライアの被害は自分の行(fc1なら784個)に閉じ込められます。
  • per-group(32):要素32個ごとにscaleを1個。被害は同じグループの31個だけに閉じ込められます。表でアウトライアを注入してもMSEがびくともしない理由です。

このper-group方式こそ、NVIDIA Blackwellがハードウェアでサポートするmicro-tensor scalingです。そして業界標準フォーマットとしてまとまったのがMXFP4:FP4要素 + 要素32個ごとの共有scale 1個。FP4という無理筋が実戦で成立する理由はここにあります。グループscaleがアウトライアの爆発半径を31要素に縮めておいたからです。

3. 増えたscaleのコスト:2^k scaleとeffective bits

もちろんタダではありません。fc1基準でscaleは1個(per-tensor)から6,272個(per-group)に増えました。これらのscaleも保存し演算すべきデータです。MXFP4はこのコストを2つのアイデアで抑えます。

第一に、scaleを2の累乗に制限します。MXFP4のグループscaleはFP32の実数ではなく、8ビット指数(E8M0)、つまり $2^k$ の形だけが許されます。こうするとscaleの乗算が指数の加算に変わり、ハードウェアが極端に安くなります。代わりにscaleは2倍単位でしか動けず、最適値からずれることがあります。その損害を実測すると:

layergroup32 + FP scalegroup32 + 2^k scale損害
fc11.23e-051.58e-051.28倍
fc22.66e-053.62e-051.36倍
fc36.12e-056.30e-051.03倍

MSEで1.0〜1.4倍です。アウトライアが作っていた22倍と比べれば、桁が違う安い保険料です。

第二に、オーバーヘッドをeffective bitsで管理します。要素1個あたりの実際の保存コストはこう計算します:

$$\text{effective bits} = \text{要素ビット} + \frac{\text{scaleビット}}{\text{グループサイズ}}$$

MXFP4は $4 + 8/32 = 4.25$ ビットです。scaleをFP16にしていたら4.5ビットだったところ、$2^k$ トリックのおかげで8ビット指数で足りるので4.25で済みます。NVIDIAのVS-Quant2は同じアイデアを階層に積みます。16個ごとに安いINT4 scaleを付け、テンソルあたり1個だけの高いFP scaleが絶対的な大きさを補正する構造で、やはり4.25ビットです。

4. FP4精度でまとめる

方式ごとにMLPのすべての重みを4ビットに量子化してMNIST精度を測りました。

FP4量子化方式別のMNIST精度

方式精度effective bits
FP32 baseline97.36%32
FP4 per-tensor97.20%4.0
FP4 per-group(32)97.37%4.5
MXFP4 (group32, 2^k)97.18%4.25

MXFP4はbaselineとの差0.2%pで7.5倍の圧縮を達成します。4ビットの重みが「無理筋」ではなく実用的な選択になる瞬間です。


残された3つの質問

ここまでがweightを4ビットに収める話でした。グループscale(処方 ①)がアウトライアの爆発半径を狭め、4ビットの重みを実用にしたわけです。しかしまだ3つの質問が残っています。

  1. これまではすべてweightの話だった。入力ごとに値が変わるactivationはどうするのか?
  2. scaleを決めた後、各値を「最も近い格子点に丸める」のは本当に最善なのか?
  3. アウトライアを閉じ込めるのではなく、消してしまうことはできないのか?

それぞれの質問に、実戦で使われる武器(処方)を1つずつで答えます。

5. 処方 ②:activationは切り捨てる — KLクリッピング

weightは学習が終われば固定されるので、min/maxを正確に知ることができます。しかしactivationは入力ごとに範囲が変わります。そこでデプロイ前に代表的な入力を数バッチ(calibrationデータ)流し、「この層のactivationはだいたいこの範囲だ」という統計を集めておく必要があります。

問題は、その統計のどの値を基準にscaleを決めるかです。観測された最大値?それではweightのときとまったく同じ病気が再発します。ReLU出力の分布は裾が非常に長いからです。実測すると、fc1のactivationの最大値は19.7ですが、値の99.9%は9.7以下にあります。最大値でscaleを決めるのは、上位0.1%のために残り99.9%の解像度を捧げることです。

ならばどこかで切り捨てる(clipping)のが得のはずですが、どこで切るべきでしょうか。早く切りすぎると切られた値の情報が消え(saturation損失)、遅く切りすぎると格子が疎になります(rounding損失)。両者の間のどこかに最適点があります。

下のアニメーションがそのトレードオフです。釣鐘型の分布の上で4ビット格子(赤い線16本)を最大値から内側へ絞っていくと、最初は格子が密になってMSEが急落しますが、分布の本体に食い込み始めると飽和損失が大きくなり、MSEは再び上昇します。

4-bit格子を絞りながら見るclippingトレードオフ — MSEがU字を描いて最適点で止まる

このU字は実際の大規模モデルでもそのまま現れます。下の図はNVIDIAのOCTAV論文3がResNet-50のweight・activation層でclipping地点をスイープしながら量子化MSEを実測した結果です。ここでのMSEはテンソルの要素ごとに測った $\mathbb{E}[(Q(x)-x)^2]$、つまり量子化前後の値の差の二乗平均です(モデル出力の差ではありません)。すべての層・ビットでU字と最適点(丸印)が見え、ビットが低いほど最適clipは内側に入ります。OCTAVはこの最適点をNewton-Raphson反復で学習の毎ステップ見つけ出す手法です。

clipping scalarに対する量子化MSEの実測 — ResNet-50のweight/activation層、4/6/8-bit

図の出典:Sakr et al., ICML 20223, Figure 1.

TensorRTの解法4はこれを情報損失の最小化問題として解きます。候補地点Tごとに「Tで切った元の分布P」と「それをnレベルに量子化した後、Pと同じ目盛りに広げ直した分布Q」を作り(解像度が違う2つの分布は直接比較できないため)、2つの分布の差であるKLダイバージェンス $D_{KL}(P\|Q)$ を計算します。この値が最小になるTが、情報を最も失わない切断点です。2つの損失が逆方向に動くので、KL曲線はU字になります。

fc1のactivation分布とKL最適クリップ地点、KL vs T曲線

我々のMLPでKLが選んだ地点はT=12.8です(最大値19.7)。これでactivationを量子化して最大値方式と比較すると:

bitsclip = maxclip = KL
497.19%97.24%
396.91%96.93%
291.21%96.15%

ビットが低くなるほど格差が開きます。2ビットではアウトライア数個を諦めた代償として5%pを回収します。この方法の実戦での強みは、分布の形に対する仮定がないことです。層ごとにactivation分布がバラバラでも(単調減少でも釣鐘型でも)同じように動作します。

6. 処方 ③:丸めを学習する — AdaRound

scaleとclipを全部決めても、自由度が1つ残っています。各値を格子点に送る丸め(rounding)です。「最も近い格子点へ」(round-to-nearest、RTN)はあまりに当然すぎて選択肢だとすら思えないのですが、QualcommのAdaRound5はこれが最適でないことを示しました。

理由はこうです。同じ層の重みは同じ入力に掛けられて1つの出力に合算されます。だから重み1つ1つの丸め誤差は独立ではなく、出力側で相殺されたり増幅されたりします。RTNは各重みの誤差だけを見て決めるので、この相互作用を丸ごと無視しています。個別最善の合計が全体最善ではないのです。

そこでAdaRoundは基準を変えます。重みではなく層の出力を最もよく復元する切り上げ/切り捨ての組み合わせを探します:

$$\arg\min_{\mathbf{V}} \|\mathbf{W}\mathbf{x} - \lfloor\lfloor\mathbf{W}\rfloor + h(\mathbf{V})\rceil\,\mathbf{x}\|_F^2 + \lambda f_{reg}(\mathbf{V})$$

数式は複雑に見えますが、構造は単純です。各重みを切り捨て($\lfloor w \rfloor$)にするか切り上げ($\lfloor w \rfloor + 1$)にするかの選択を $h(V) \in (0,1)$ という連続値に緩和してgradientで学習し、正則化項 $f_{reg}$ が学習の終わりにこの値を0か1に押し込みます。ラベルも全体の再学習も不要です。calibration入力数バッチと層ごとの短い最適化で済みます。ノートブックでは層あたり800ステップ、CPUで数十秒でした。

RTN vs AdaRound — INT3/INT2精度

bitsRTNAdaRound丸めが変わった割合
INT396.83%97.23%11.8%
INT253.37%96.56%12.0%

INT2が劇的です。丸めの決定の12%を変えただけで、53%が97%に戻ります。「最も近い値」という直感が、低ビットでいかに高くつく直感だったかを示しています。

7. 処方 ④:座標系を回してアウトライアを消す — Hadamard回転

これまでのすべての処方はアウトライアとの共存でした。閉じ込める(グループscale)、切り捨てる(clipping)、避けて回る(AdaRound)。最新のLLM量子化(QuaRot6、SpinQuant7)は発想が違います。アウトライアが存在しない座標系に回転してしまうのです。

これが可能なのは計算不変性のおかげです。直交行列 $R$($RR^\top = I$)を1つ選ぶと、

$$\mathbf{y} = \mathbf{W}\mathbf{x} = (\mathbf{W}R^\top)(R\,\mathbf{x})$$

重みを $WR^\top$ に、入力を $Rx$ に事前にすり替えても、出力は数学的に完全に同一です。モデルが計算する関数はそのままで、量子化から見えるテンソルの「形」だけが変わるのです。

ここで $R$ にHadamard行列(成分がすべて±1の直交行列を $\sqrt{n}$ で割ったもの)を選ぶと、魔法が起きます。回転後の各成分は、元の全成分の±平均になります。つまり一箇所に集中していたアウトライアが、全次元に $1/\sqrt{n}$ の大きさで薄く塗り広げられるのです。尖っていた分布が釣鐘型に戻ります。

下のウィジェットで試してみてください。[1, 0] のように一軸に集中したベクトルをスライダーで回転させると、45°で2つの座標が [0.707, 0.707] と等しくなります — 長さはそのままで、アウトライアが2軸に均等に分けられます。

前のアウトライアシナリオ(重み1個を50倍に)をfc2で再現し、回転してみました:

Hadamard回転前後の重み分布 — アウトライアが消えて釣鐘型に

分布がどれだけ「正常」かを測る指標である尖度(kurtosis、ガウシアンなら3)で見ると、13,717から54へ、入力・出力の両側を回転(QuaRot方式)すると3.4まで下がります。アウトライアの痕跡が統計的に消滅したのです。FP4で量子化した後の出力誤差で確認すると:

per-tensorper-group(32)
アウトライアなし(参照)0.0590.045
アウトライア + 回転なし0.8590.043
アウトライア + Hadamard 入力側のみ0.1970.089
アウトライア + Hadamard 両側(QuaRot)0.0930.067

読み方はこうです。アウトライアがあるとper-tensorは全滅しますが(0.86)、両側回転が参照水準の近く(0.09)まで蘇らせます。逆にper-groupは回転しても利得がありません。グループscaleがすでにアウトライアを閉じ込めていたからです。ここに回転の性格が現れています。回転の利得はscaleが粗いほど大きい。だから回転はscaleを細かく付けにくい場所、代表的にはW4A4のactivationで真価を発揮します。

コスト面でも魅力的です。グループscaleには保存オーバーヘッド(+0.25ビット)が付きますが、回転は追加保存が0です。$WR^\top$ はデプロイ前に事前計算しておけばよく、推論中の $Rx$ はHadamard変換の構造のおかげで $O(n \log n)$ で処理できます。

1つ注意点があります。我々の実験は256次元なので、アウトライアは1/16にしか縮みませんでした。実際のLLMはhidden次元が4096以上なので、1/64以下に縮みます。次元が大きいほど回転はより完璧になります。QuaRotがLLaMA-2 70BのW4A4推論を成立させた核心のトリックがこれで、SpinQuantは回転行列そのものを学習してさらに一歩進みます。


自分で回してみる

上のplotと表は下のノートブックを実際に回して出たものです。ランタイムからすべて実行で終わります。

  • 📓 low-bit-quantization.ipynb — FP4/MXFP4グループ量子化、アウトライア実験、KLクリッピング、AdaRound、Hadamard回転

まとめ — 4つの処方箋を1つの表に

低ビット量子化は「平均的な値」ではなく「最悪の値(アウトライア)」との戦いです。scaleは常に最大値に縛られるからです。その戦いの処方を4つ、1つの表にまとめると:

処方アイデアコスト我々の実測
① グループscale(MXFP4)被害をグループの31個に閉じ込める+0.25 bitFP4でbaseline −0.2%p
② KLクリッピングアウトライア数個を諦めて解像度を守るcalibration探索2-bitで+5%p
③ AdaRound丸めを層出力基準で学習する層ごとの短い最適化INT2で53→97%
④ Hadamard回転アウトライアのない座標系へ回転する追加保存0per-tensor誤差9倍減
  • 4つの処方は競合ではなく直交するツールです。実際のW4A4パイプラインは回転(アウトライア除去)+ グループscale(残余分散の吸収)+ クリッピング(activation)を併用します。

References


  1. Song Han. MIT 6.5940 TinyML and Efficient Deep Learning Computing, Lecture 6: Quantization Part II. efficientml.ai ↩︎

  2. Dai et al. VS-Quant: Per-Vector Scaled Quantization for Accurate Low-Precision Neural Network Inference. MLSys 2021. ↩︎

  3. Sakr et al. Optimal Clipping and Magnitude-aware Differentiation for Improved Quantization-aware Training. ICML 2022. ↩︎ ↩︎

  4. Szymon Migacz. 8-bit Inference with TensorRT. GTC 2017. ↩︎

  5. Nagel et al. Up or Down? Adaptive Rounding for Post-Training Quantization. PMLR 2020. ↩︎

  6. Ashkboos et al. QuaRot: Outlier-Free 4-Bit Inference in Rotated LLMs. NeurIPS 2024. ↩︎

  7. Liu et al. SpinQuant: LLM Quantization with Learned Rotations. 2024. ↩︎