TT Lab
はじめる
学ぶ 学習パス コース

Transformer — アテンションを手で計算する

int8 に畳んで戻すと何がどれだけ変わるか

TT Labで続きを見る

一言でいうと

量子化は、実数をスケール1つと整数1つに変える作業で、失うものはスケールが決めます。そのスケールを決めるのは、リストの中で最も大きい絶対値1つです。

なぜ必要なのか

モデルをint8に変えると、重みが4分の1に減り、整数の乗算器を使えます。そのため、誰もが一度は試します。問題はそのあとです。精度が少し落ちるのに、どこで落ちたのかを言えません。

ツールは、1行で済みます。量子化の関数を呼べば、モデルが出てきます。その中でどんな演算が起きたのかを知らないと、できることは、オプションを変えながら再び動かしてみることだけで、それは直すことではなく、運を試すことです。

そのため、この文章と次のラボは、ツールを使いません。値8個のリスト1つを、手で畳んでは広げてみることから始めます。そこで見えることは、大きなモデルでもそのまま見えます。

「畳む」とはどういう意味か

実数のリスト1つをint8に移すには、2つのことを決める必要があります。スケール(scale)とゼロ点(zero point)です。

最も単純なものが、対称量子化です。スケールをmax(|x|) / 127にして、値をスケールで割って丸めます。最も大きい絶対値がコード127に届き、0はコード0にそのまま収まります。

scale = max(abs(x) for x in values) / 127
codes = [round(x / scale) for x in values]   # -127 부터 127 까지
back  = [c * scale for c in codes]           # 편 값. 원본이 아니다

非対称(アフィン)量子化は、値が片側に偏っているときに使います。幅を(max - min) / 255にして、256段階をすべて使い、実数0がどのコードに収まるかを、ゼロ点として別に持ち歩きます。復元は(code - zero_point) * scaleです。整数演算だけで推論する方法をまとめた論文が、この式をそのまま使っています。

どちらの方式でも、失うものは同じ場所にあります。スケールの半分です。スケールが0.007なら、どんな値でも最大で0.0035だけずれます。値が大きくても小さくても、同じだけです。

まず丸めの規則を固める

ここで、人が最初に引っかかります。Pythonのroundは、学校で習った四捨五入ではありません。

round(0.5)   # 0   — 1 이 아니다
round(1.5)   # 2
round(2.5)   # 2   — 3 이 아니다

ちょうど0.5の位置を、偶数のほうへ送ります。0.5をいつも上に切り上げると、丸めた値の平均が少しずつ上にずれるので、このように決まっています。量子化は、配列全体に丸めを1回ずつ行う作業なので、このずれがそのままモデルの偏りになります。

問題は規則ではなく、規則が2通りあるという事実です。ある実装は偶数に送り、ある実装は0から遠いほうに送ります。同じ重みを同じスケールで畳んだのに、コードが1つずつ違っていると、そのあとのすべての比較が意味を失います。そのため、量子化のコードを読むときに、スケールの式より先に確認することは、丸めの規則です。

スケールを決めるのは値1つ

対称量子化のスケールは、max(|x|) / 127です。この式には、平均も分散もありません。最も大きい絶対値1つがすべてです。

そのため、リストに際立って大きい値が1つ混ざると、その1つが残りすべての精度を決めます。残りがすべて1未満なのに、1つが42なら、スケールが42/127になり、1未満の値は、コード0から3の間につぶれます。256段階があるのに、4段階しか使わないのです。

値8個のうち、42の1つがスケールを決める図。行列全体にスケール1つを使うと、1未満の7つの値が、コード0から3までの4段階につぶれ、行ごとにスケールを別に設定すると、42がない行は、コード28から127まで、目盛りをきちんと使い分けます

これが、LLM.int8()論文の出発点です。大きな言語モデルの隠れ状態には、他の値よりはるかに大きい成分が現れますが、その数個のために、残りのすべてが使えなくなるという観察でした。解決の方向も、そこから出てきます。スケールをより小さい単位で設定するか、大きいものを別に取り出すかです。

スケールをより小さい単位で設定するほうが、先にすることです。行列全体をスケール1つで畳む代わりに、行ごと(または列ごと)にスケールを別々に設定すれば、小さい行が大きい行に引きずられません。ただし、無料ではありません。スケールを行ごとに持ち歩く必要があり、行列の積が成り立つには、スケールが行か列の単位である必要があります。値ごとにスケールが違うと、整数の累積の中からスケールを取り出せません。

整数で掛けると何が残るのか

整数推論の核心は、乗算と累積がすべて整数であることです。コード同士を掛けて足す間は、丸めが一度も起きません。すべて足したあとで、最後にスケール2つを掛けて、実数に広げます。

acc  = sum(a_code[k] * b_code[k] for k in range(d))   # 정수만
value = acc * a_scale * b_scale                       # 마지막에 한 번

そのため、行列の積で生じる誤差は、累積で膨らんだものではなく、最初に畳むときにすでに生じたものです。原因を探す場所が1か所だけだということで、これはよい知らせです。

残る問いは、その誤差があとでどうなるのかです。アテンションのスコアは、ソフトマックスを通って確率になります。ソフトマックスは、差を指数で広げるので、スコアの小さな誤差が確率で大きくなることもあり、逆に埋もれることもあります。どちらなのかは、測ってみる前には言えません。次のラボで、自分で測ります。

現場での姿

第1に、精度が少し落ちたのに、どこで落ちたのかがわかりません。ツールが1行なので、内側を見る目がありません。スケール・丸め・単位のうち、何を変えたときに何が動くのかを、手で一度やってみた人だけが、指摘できます。

第2に、同じモデルを2つのツールで量子化したら、結果が違います。スケールの式が同じでも、丸めの規則が違うと、コードが1つずつずれます。どちらが正しいかではなく、何が違うのかを先に確認する必要があります。

第3に、テンソル単位で畳んだら、特定の層だけが崩れます。その層の重みの分布に、際立って大きい値がある場合です。全体の平均誤差は問題なさそうに見えるのに、小さい行の相対誤差だけが爆発します。

第4に、活性値(activation)は、重みよりはるかに厄介です。重みは固定なので、一度測れば終わりですが、活性値は入力ごとに変わります。キャリブレーション(calibration)用のデータで範囲を測りますが、そのデータが実際の入力と違うと、スケールがずれます。

第5に、サイズと速度を同じものとして話します。ファイルが4分の1になったことと、実際に速くなったことは、別のことです。整数カーネルが実際に動いたかどうかは、別に確認する必要があります。

実務で本当に大切なこと

次のラボですること

/root/work/tf-quant/quant.pyを、1ステップずつ育てていきます。ツールを呼ばず、演算を自分で行います。このPodのシステムのPythonにはnumpyがありません(numpyは/opt/onnx-labの中にしかありません)。そのため、リストとリストのリストを、標準ライブラリで扱います。Pythonのmathモジュールのfloor・expで十分です。

丸めの規則を2通り作って、どこで分かれるかを確認することから始めます。そのあと、対称・非対称の量子化を作り、畳んでから広げた値が、元からどれだけ離れたかを測る物差しを作ります。ここまでが半分です。

残りの半分が、このラボの要点です。ふつうの値に大きい値を1つ付け足して、残りの値が何段階に減るかを数えます。行ごとに幅が大きく異なる行列を置いて、テンソル単位と行単位を並べて測ってみます。コード同士を整数だけで掛けて、累積がちょうど合うかを確認し、最後に、その誤差がソフトマックスを通ったあとの確率でどれだけになるかを測ります。外れ値を仕込んだキーで、もう一度測って、前に見たことがアテンションでどう現れるかを見ます。

採点ツールは、書いておいた説明を信用しません。自分で作ったモジュールを実際に呼び出し、毎回異なる入力で関数を直接叩いて、採点ツールが別に計算した値と照合します。照合の大部分は、整数の配列同士です。コード1つがずれれば、すぐに現れます。