この記事は軽量化シリーズの第5回(#5)です。低ビット量子化 #4が「学習し終えたモデルを事後にうまく量子化する(PTQ)」処方箋だったなら、今回はいっそ学習の段階に量子化を組み込むQATです。コードは github.com/warpspaceinc/efficient-ml-practiceqat.ipynb にあり、以下のすべてのplotと表はそのコードを実際に回して測定した我々の数値です。

PTQが崩れる地点

これまでの量子化はすべて PTQ(Post-Training Quantization) — 学習し終えたFP32モデルを事後に整数格子へ移すものでした。大きなモデルにはよく効きます。INT8 PTQでResNet-50は−0.1%、GoogleNetはほぼ無損失です。

問題は小さなモデルです。同じINT8 PTQなのに数値が印象的です1。例えば:

モデルFP32INT8 PTQ (per-tensor)
ResNet-5076.1%−0.1%
MobileNetV170.9%0.1%
MobileNetV271.9%0.1%

MobileNetは 70.9%から0.1%へ、つまり完全に死にます。パラメータを切り詰めて作ったコンパクトなモデルほど余裕(redundancy)がなく、量子化誤差を吸収する場所がないからです。そしてこの現象はビットを下げるほど(4ビット以下)モデルサイズと無関係に現れます。

我々のMNIST MLP(784→256→128→10)でも同じことが再現されます。重みをビット幅ごとにPTQしたとき:

  • 8ビット:97.1%(無事)
  • 4ビット:96.9%
  • 3ビット:94.9%
  • 2ビット:13.6% — 事実上ランダム(10%)に近く崩壊

ではどうすればいいのか。答えはシンプルでいて強力です。「どうせ量子化された状態で推論するなら、学習のときから量子化された状態で学習しよう。」


QAT:学習中に量子化を真似る

QAT(Quantization-Aware Training)の核心は fake quantization(模擬量子化) です。学習のforward passに、実際の推論で起きる量子化をあらかじめシミュレートして差し込みます。構造は3つの原則にまとめられます2

  1. FP32のマスター重み $W$ を維持し続ける。 学習される実体はこの連続値です。
  2. forwardでのみ、重みと活性値を量子化してから復元(fake-quant)して、整数格子上の値で演算します。
  3. 推論時にはこの量子化された重みだけを使います。

数式で見ると、#2のアフィン写像をそのままforwardに挿入することです。重みと出力(活性値)それぞれに:

$$Q(W) = S_W\, q_W, \qquad Q(Y) = S_Y\,(q_Y - Z_Y)$$

ここで $q_W = \mathrm{round}(W/S_W)$ のように roundで格子にスナップしたあと、再びscaleを掛けて実数に戻します(だから「fake」— 値は格子上にありますが、演算自体はFP32で行います)。学習がこの格子の上で進むので、モデルは「この重みを2ビットに潰しても出力が良くなるように」自らを調整するようになります。量子化誤差を事後に耐えるのではなく、学習目標に最初から織り込むわけです。

ところがここで決定的な問題が一つ生じます。

STE:roundは微分が0なのに、どう学習するのか

fake quantizationの心臓は round 関数です。そしてroundは階段関数です — ほとんどすべての点で傾きが0で、格子境界でのみ無限大です。

$$\frac{\partial Q(W)}{\partial W} = 0 \quad (\text{ほとんどいたるところで})$$

逆伝播の連鎖律をそのままたどると惨事です。

$$g_W = \frac{\partial L}{\partial W} = \frac{\partial L}{\partial Q(W)} \cdot \frac{\partial Q(W)}{\partial W} = \frac{\partial L}{\partial Q(W)} \cdot 0 = 0$$

すべての重みの勾配が0になります。学習が一歩も進めません。round一つが逆伝播を丸ごと塞いでしまうのです。

解決策が STE(Straight-Through Estimator) です34。発想は図々しいほど単純です。「forwardではroundを使うが、backwardではroundが無かったことにして勾配をそのまま通す。」 つまり量子化関数を微分するとき、それを恒等関数(identity)とみなします。

$$g_W = \frac{\partial L}{\partial W} \;\approx\; \frac{\partial L}{\partial Q(W)}$$

forwardは階段、backwardは傾き1。この「嘘」がQATを成立させます。

面白いのはこの図々しいトリックの出どころです。STEはディープラーニングの父 ジェフリー・ヒントン(Geoffrey Hinton) が2012年のCoursera講義で通りすがりに紹介したアイデアで3、翌年にBengioらが定式化しました4。そのヒントンはチューリング賞(2018)に続き 2024年のノーベル物理学賞まで受賞し、最近までAIを語りながら精力的に活動しています。10年余り前の講義でぽろりと放った一行が、今も低ビット学習の根幹を支えているわけです。

STE — forwardは階段関数、backwardは恒等関数(傾き1)として通過

左がforwardです。$Q(w)$ は点線(恒等関数 $y=w$)の周りを階段で近似します。右がbackwardです。真の微分は0(赤)ですが、STEは表現範囲の中で勾配を1(緑)として通します。範囲の外(clipされた領域)では0にして、格子を外れた値には信号を与えません。

下で直接触ってみてください。スカラー重み $w$ 一つ分の計算グラフ($w \to Q(w) \to \hat y=Q(w)\cdot x \to L=\tfrac12(\hat y-t)^2$)です。重み $w$・入力 $x$・正解 $t$ をスライダーで変えると、lossと勾配がリアルタイムに変わります — 特にbackwardで真の微分(0)とSTE勾配がどう分かれるかを見てください。

PyTorchではカスタム autograd.Function 一つで済みます。forwardに量子化、backwardに通過(+clipマスク)を入れます。

class FakeQuantSTE(torch.autograd.Function):
    @staticmethod
    def forward(ctx, w, n_bits):
        qmax = 2 ** (n_bits - 1) - 1
        S = w.detach().abs().max() / qmax + 1e-12       # per-tensor symmetric scale
        q = torch.clamp(torch.round(w / S), -qmax, qmax)  # 格子にスナップ
        ctx.save_for_backward((w.abs() <= S * qmax).to(w.dtype))  # clipマスク
        return q * S                                     # 復元(fake-quant)
    @staticmethod
    def backward(ctx, g):
        (mask,) = ctx.saved_tensors
        return g * mask, None          # STE: 範囲内は通過、外は0

そしてQAT学習は、forwardでこの fq を重みに被せたまま普通に回せば済みます。マスター重みはFP32のまま残り、STEのおかげで勾配がそのマスターへ流れていきます。

def qat_forward(x):
    x = F.relu(F.linear(x, fq(fc1.weight, b), fc1.bias))
    x = F.relu(F.linear(x, fq(fc2.weight, b), fc2.bias))
    return F.linear(x, fq(fc3.weight, b), fc3.bias)

直接測定:PTQ vs QAT

同じMNIST MLPを、同じビット幅で PTQ(事後量子化) と QAT(FP32モデルからfake-quantで2エポックfine-tune) でそれぞれ測り直しました。

ビット幅別のPTQ vs QAT精度 — 2ビットでPTQは崩壊、QATは回復

bitsPTQQATFP32
897.1%97.9%97.1%
496.9%97.7%97.1%
394.9%97.8%97.1%
213.6%91.9%97.1%

読むべき図は明確です。

  • 8ビットでは両方とも無損失。 余裕のある区間ではわざわざQATは要りません。
  • ビットが下がるほど差が爆発。 2ビットでPTQは13.6%(ランダム水準)で死ぬのに、QATは 91.9% まで蘇らせます。MobileNetV1(INT8 PTQ 0.1%からQAT 70%台へ)とまったく同じパターンです1
  • QATがbaselineをわずかに超えることもあり(3・4ビット)、fine-tuneで数エポック多く学習された効果です。

核心はこれです。PTQは「すでに決まった重みを格子にできるだけよく合わせる」問題でした。QATは「格子の上で最も良い重みを最初から探す」問題です。自由度が違うので、低ビットで結果が分かれます。


自分で回してみる

上のplotと表は、下のノートブックを直接回して出たものです。ランタイム → すべて実行 で終わりです。

  • 📓 qat.ipynb — fake quantization、STE(autograd.Function)、PTQ vs QAT比較

STEの次:scaleとclipも学習する — LSQ・PACT

これまで我々のfake-quantはscale $S$ を $|W|_{\max}/q_{\max}$ という固定ヒューリスティックで決めていました。STEは重みの勾配を流すだけで、肝心の格子の間隔 $S$ そのものは学習されませんでした。次の改善は、そのscaleとclipまで学習することです。

LSQ — step sizeを学習する。 LSQ(Learned Step Size Quantization)5は $S$ を学習パラメータとして置きます。$Q(w)=\mathrm{round}(\mathrm{clip}(w/S))\cdot S$ で $\partial Q/\partial S$ を量子化状態遷移に敏感に計算し、勾配で最適化します — 範囲内では $\partial Q/\partial S = \mathrm{round}(w/S) - w/S$(格子値と連続値の差)、範囲外では $\pm q_{\max}$ です。効果は低ビットで大きい。我々のMLPを固定scale QATと比べると:

bitsQAT(固定scale)QAT + LSQ(scale学習)
497.9%97.7%
397.6%97.9%
292.3%96.7%

2ビットで 92%から97%へ。格子の間隔一つをデータに決めさせただけで、最後の差が埋まります。余裕のない低ビットほど「scaleをいくつにするか」が精度を分けます。

PACT — activationのclip範囲を学習する。 重みは学習が終わればmin/maxが固定ですが、activation(ReLU出力)は範囲が開いています。#4でKLで切り点を選んだなら、PACT6はそのclipしきい値 $\alpha$ を学習パラメータとして置きます。ReLUを $\mathrm{clip}(x, 0, \alpha)$ に置き換えて $\alpha$ を勾配で学習し、activation量子化の範囲をデータに決めさせます。LSQが重みのscaleを学ぶように、PACTはactivationのclipを学ぶわけです。

そしてLLM時代に入ると、この話ははるかに先まで進みます — 三値 {−1, 0, +1} で最初から学習するBitNet、FP4でプリトレーニングする流れまで。それは別のLLM量子化編で改めて扱います。

まとめ

  • PTQは大きなモデルには効くが、小さなモデル・低ビットでは崩れる。 MobileNet INT8が0.1%で死に、我々のMLPも2ビットPTQで13.6%まで崩壊する。
  • QATは学習のforwardにfake quantizationを差し込み、モデルが量子化された状態で良くなるよう自らを調整させる。FP32のマスター重みは維持される。
  • roundの勾配は0なので素朴には学習できないが、STEがbackwardで量子化を恒等関数とみなし、勾配を通す。
  • 結果:2ビットでPTQ 13.6%からQAT 91.9%へ。低ビットほどQATの値打ちが大きくなる。

これで量子化編(#2#4・#5)が締めくくられます。データ型から量子化、プルーニング、低ビット、QATまで、モデルを小さく速くする道具を概念から実測まで見渡しました。

References


  1. Krishnamoorthi. Quantizing Deep Convolutional Networks for Efficient Inference: A Whitepaper. arXiv 2018.(MobileNet PTQ・QAT数値の原典) ↩︎ ↩︎

  2. Jacob et al. Quantization and Training of Neural Networks for Efficient Integer-Arithmetic-Only Inference. CVPR 2018. ↩︎

  3. Hinton et al. Neural Networks for Machine Learning. Coursera Lecture, 2012.(STEの起源) ↩︎ ↩︎

  4. Bengio et al. Estimating or Propagating Gradients Through Stochastic Neurons for Conditional Computation. arXiv 2013. ↩︎ ↩︎

  5. Esser et al. Learned Step Size Quantization. ICLR 2020. ↩︎

  6. Choi et al. PACT: Parameterized Clipping Activation for Quantized Neural Networks. arXiv 2018. ↩︎