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

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

int8 の算術を手で解く

TT Labで続きを見る

目標

整数量子化の演算を、標準ライブラリだけで自分で行います。丸めの規則を固めることから始めて、対称・非対称の量子化を作り、畳んでから広げた値の誤差を測り、外れ値1つがスケールをどれだけ壊すかを数え、テンソル単位と行単位を並べて比べ、コード同士を整数だけで掛け、最後に、その誤差がソフトマックスを通ったあと、確率でどれだけになるかを測ります。

なぜ重要なのか

量子化のツールは、1行です。その中でどんな演算が起きたのかを知らないと、精度が落ちたときに、オプションを変えながら再び動かしてみるくらいしかできません。スケールが何で決まるのか、丸めがどちらへ行くのか、単位を絞ると何が変わるのかを、値8個のリストで一度見ておけば、大きなモデルでも、同じ箇所を指摘できます。 このラボは、ツールを使いません。このPodのシステムのPythonにはnumpyがなく(numpyは/opt/onnx-lab/bin/pythonの中にしかありません)、モデルも呼び出しません。そのため、「あるモデルはint8で精度が何パーセント落ちる」といった話は、ここではしません。自分が作ったリストと行列から測った数字だけを使います。 難しいのは式ではなく、細部です。0.5をどちらへ送るか、範囲外をクリップするかどうか、スケールを何の最大値にするかが決まっていないと、同じ入力でも、コードが1つずつずれます。 採点ツールは、書いておいた説明を信用しません。自分で作ったモジュールを実際に呼び出し、毎回異なる入力で関数を直接叩いて、採点ツールが別に計算した値と照合します。照合の大部分は、整数の配列同士なので、揺らぎがありません。

ステップ

  1. /root/work/tf-quant/quant.pyにQMAX = 127・UMAX = 255と、round_half_even(x)・round_half_away(x)・rounding_gap(values)を作成してください。2つの丸めの規則がどこで分かれるかを、目で確認します。
  2. sym_scale(values)・quantize_sym(values, scale)・dequantize_sym(codes, scale)を追加して、対称量子化を作成してください。スケールはmax(|x|) / 127です。
  3. affine_params(values)・quantize_affine(values, scale, zero_point)・dequantize_affine(codes, scale, zero_point)を追加して、非対称量子化を作成してください。スケールは(max - min) / 255です。
  4. levels_used(codes)とerror_stats(original, restored)を追加して、誤差を測る物差しを作成してください。
  5. outlier_effect(values, outlier)を追加して、大きな値1つが残りの値に何をするかを測るようにしてください。
  6. quantize_tensor(matrix)・quantize_rows(matrix)・granularity_gap(matrix)を追加して、テンソル単位と行単位を比べてください。
  7. transpose(matrix)・int_matmul(left_codes, right_codes)・float_matmul(left, right)・quant_matmul(left, right)を追加して、整数だけで行列を掛けてください。
  8. WEIGHTS・OUTLIER・QUERIES・KEYSと、softmax(scores)・attention_shift(queries, keys)を作成し、/root/work/tf-quant/quant_report.jsonと、/root/work/tf-quant/quant_report.mdに結果を記録してください。

参考

丸めの規則を固める

/root/work/tf-quant/quant.pyにQMAX = 127・UMAX = 255と、round_half_even(x)・round_half_away(x)・rounding_gap(values)を作成してください。前者は、Pythonの標準のroundそのままに、ちょうど0.5を偶数のほうへ送り、後者は、0から遠いほうへ送ります。rounding_gapは、2つの規則が分かれる値だけを、受け取った順序のまま集めて返します。

まずpython3 -c "print(round(0.5), round(1.5), round(2.5))"を動かしてみてください。0 2 2と出ます。round_half_awayをmath.floor(x + 0.5)の1行で書くと、負の数で間違います。-1.5は-2へ行く必要がありますが、その式は-1を返します。math.floor(x)で下を求め、小数部を別に見て、0.5の位置だけを符号で分けてください。2つの関数は、どちらも整数を返します。

対称量子化で畳んで広げる

sym_scale(values)・quantize_sym(values, scale)・dequantize_sym(codes, scale)を追加してください。スケールはmax(|x|) / QMAXで、値がすべて0なら1.0です。畳むときは、スケールで割ってround_half_evenで丸めたあと、-QMAXからQMAXまでにクリップします。

スケールを最大値ではなく最大の絶対値にしなければ、負の側も範囲に収まりません。クリップは、max(-QMAX, min(QMAX, code))の1行で済みます。広げる関数は、スケールを掛けるだけなので、1行です。広げた値が元と同じにならないのが正常です。失った分は、スケールの半分の範囲に収まります。

非対称量子化で256段階をすべて使う

affine_params(values)・quantize_affine(values, scale, zero_point)・dequantize_affine(codes, scale, zero_point)を追加してください。スケールは(max - min) / UMAX、ゼロ点はround_half_even(-min / scale)を0からUMAXまでにクリップした値です。最大と最小が同じなら、(1.0, 0)です。畳むときは、round_half_even(x / scale) + zero_pointを、0からUMAXまでにクリップします。

クリップを抜かすと、ここで実際に壊れます。両端が丸めで1つずつずれると、round(x/scale) + zero_pointが256になって、uint8の範囲を超えます。広げる式は、(code - zero_point) * scaleです。この式のおかげで、実数0は誤差なく元に戻ります。ゼロ点がある理由が、それです。

ずれ幅を測る物差しを作る

levels_used(codes)とerror_stats(original, restored)を追加してください。前者は、異なるコードの個数で、後者は、max_abs・mean_abs・max_relの3つのキーを持つ辞書です。max_relは、最大の絶対誤差を元の最大の絶対値で割った値で、その最大の絶対値が0なら0.0です。

相対誤差を値1つ1つで割ると、0の近くの値で無限大に跳ね上がります。そのため、リストが持っている幅で割ります。スケールがその幅で決まるので、比べる対象もそれです。levels_usedは、len(set(codes))の1行です。256段階があるのに、いくつの段階しか使っていないかが、次のステップの話です。

大きな値1つが残りに及ぼす影響を測る

outlier_effect(values, outlier)を追加してください。valuesだけを畳んだときと、values + [outlier]を畳んだときを、それぞれ測り、もとのvaluesの部分だけを比べます。返すキーは、clean_scale・dirty_scale・clean_levels・dirty_levels・clean_max_abs・dirty_max_absの6つです。

外れ値を付けたリストをまるごと畳んだあと、先頭からlen(values)個だけを切り出して比べてください。外れ値自身の誤差まで数えると、話が逆になります。壊れるのは、外れ値ではなく、その隣にいたふつうの値たちです。dirty_levelsがいくつに落ちるかを、目で見てください。256段階があるのに、いくつの段階を使っていますか。

テンソル1つのスケールと行ごとのスケールを比べる

quantize_tensor(matrix)・quantize_rows(matrix)・granularity_gap(matrix)を追加してください。前の2つは、それぞれ(배율, 코드 행렬)(プレースホルダーはスケールとコード行列です)と、(배율 목록, 코드 행렬)(プレースホルダーはスケールのリストとコード行列です)を返します。granularity_gapのキーは、tensor_max_abs・row_max_abs・tensor_worst_row_rel・row_worst_row_rel・tensor_worst_row_levels・row_worst_row_levelsの6つです。

worst_row_relは、行ごとにその行のerror_statsでmax_relを求めたあとの最大値で、worst_row_levelsは、行ごとにlevels_usedを求めたあとの最小値です。最大の絶対誤差だけを見ると、2つの方式の差はほとんど見えません。その値は、最も大きい行が決めるからです。小さい行がどうつぶれるかは、相対誤差と、使った段階の数に現れます。

整数だけで行列を掛ける

transpose(matrix)・int_matmul(left_codes, right_codes)・float_matmul(left, right)・quant_matmul(left, right)を追加してください。int_matmulは、積も和もすべて整数である必要があります。quant_matmulは、左を行単位で、右を列単位で畳んだあと、整数で掛けて、acc * left_scale * right_scaleで広げ、(복원 행렬, 정수 누적 행렬, 왼쪽 배율 목록, 오른쪽 배율 목록)(プレースホルダーは復元行列、整数の累積行列、左のスケールのリスト、右のスケールのリストです)を返します。

右を列単位で畳むには、transposeで列を取り出し、各列にquantize_symを使って、またtransposeで戻せば済みます。累積の中では、スケールを絶対に掛けないでください。最後に1回だけ掛けます。そのため、行列の積の誤差は、累積で膨らんだものではなく、最初に畳むときにすでに生じたものです。採点ツールは、整数の累積行列をそのまま照合するので、1つずれただけでも捕まります。

ソフトマックスを通ったあとに残るものを測る

WEIGHTS(6行8列)・OUTLIER(絶対値が20.0以上)・QUERIES(4行8列)・KEYS(5行8列)と、softmax(scores)・attention_shift(queries, keys)を作成してください。WEIGHTSは、最も大きい行の最大の絶対値が、最も小さい行の10倍以上である必要があり、3つの行列のすべての値は、絶対値が2.0以下です。そのあと、/root/work/tf-quant/quant_report.jsonにはsym_scale・sym_max_abs・sym_mean_abs・sym_max_rel・sym_levels・affine_scale・affine_zero_point・affine_max_abs・affine_levels・outlier_clean_scale・outlier_dirty_scale・outlier_clean_levels・outlier_dirty_levels・outlier_clean_max_abs・outlier_dirty_max_abs・tensor_max_abs・row_max_abs・tensor_worst_row_rel・row_worst_row_rel・tensor_worst_row_levels・row_worst_row_levels・matmul_max_abs・matmul_max_rel・clean_score_max_abs・clean_prob_max_abs・clean_argmax_changed・dirty_score_max_abs・dirty_prob_max_abs・dirty_argmax_changedを、/root/work/tf-quant/quant_report.mdには## 무엇을 쟀나 ## 대칭과 비대칭은 어디서 갈렸나 ## 이상치 하나가 한 일 ## 행렬 곱과 소프트맥스를 지나면の4つの節で書いてください(韓国語の見出しは順に、「何を測ったか」「対称と非対称はどこで分かれたか」「外れ値1つがしたこと」「行列の積とソフトマックスを通ると」という意味です)。

数字は手で書かず、自分のコードを実際に動かして得た値で埋めてください。sym_*とaffine_*は、WEIGHTSを1行に広げたリストについての値で、outlier_*は、そのリストとOUTLIERでoutlier_effectを呼び出した結果です。tensor_*・row_*は、granularity_gap(WEIGHTS)の値で、matmul_*は、quant_matmul(QUERIES, transpose(KEYS))の復元行列を、float_matmul(QUERIES, transpose(KEYS))と比べたerror_statsのmax_abs・max_relです。clean_*はattention_shift(QUERIES, KEYS)、dirty_*は、KEYSをコピーして[0][0]の位置だけをOUTLIERに変えた行列で呼び出した結果です。元のKEYSは書き換えないでください。スコアの誤差と確率の誤差がどちらへ動くかを見て、見たとおりに書いてください。推測して書かないでください。