tune_gemm の設計

設計目標

tune_gemm は tune のための参照ドメインです。ブロッキングのパラメータが重要で、正しさを厳密に検査でき、すべての候補をシリアライズできる問題です。ハードウェア固有のカーネルに依存せずに、行列積の上でチューニングの規律全体(計時の前に検証する、ワークスペースを考慮する、端数をテストする、設定と環境を記録する)を示します。

数学的背景

積とそのレイアウト

各バッチ bb について Cb=AbBbC_b = A_b B_b であり、cij=∑l=0k−1ail bljc_{ij} = \sum_{l=0}^{k-1} a_{il}\, b_{lj} で、2mnk2mnk 回の浮動小数点演算がかかります。したがって 1 操作あたり tt µs の計測時間は次に対応します

GFLOP/s=2 m n k⋅batcht⋅103.\text{GFLOP/s} = \frac{2\,m\,n\,k\cdot\text{batch}}{t \cdot 10^{3}} .

r×cr \times c 行列は、バッチ bb の要素 (i,j)(i, j) を、行優先のレイアウトでは b rc+i c+jb\,rc + i\,c + j に、列優先のレイアウトでは b rc+j r+ib\,rc + j\,r + i に格納します。どちらの写像もインデックスの 3 つ組から [0,batch⋅rc)[0, \text{batch}\cdot rc) への全単射なので、どのレイアウトの組み合わせも同じ数学的な積を表します。

ブロッキングとワークスペース

gemm_blocked は反復空間を mc×nc×kcmc \times nc \times kc のブロックに分割し、各ブロックを mr×nrmr \times nr のレジスタタイルに分割します。1 つのブロックは AA の mc×kcmc \times kc ブロック、BB の kc×nckc \times nc パネル、CC の mc×ncmc \times nc ブロックに触れます:

W=8 (mc⋅kc+kc⋅nc+mc⋅nc) bytes,I=2 mc nc kcW=mc nc kc4 (mc⋅kc+kc⋅nc+mc⋅nc) flop/byte.W = 8\,(mc\cdot kc + kc\cdot nc + mc\cdot nc) \ \text{bytes}, \qquad I = \frac{2\, mc\, nc\, kc}{W} = \frac{mc\, nc\, kc}{4\,(mc\cdot kc + kc\cdot nc + mc\cdot nc)} \ \text{flop/byte}.

立方体 mc=nc=kc=smc = nc = kc = s では W=24s2W = 24 s^2、I=s/12I = s/12 です。ブロックが大きいほど、読み込んだ各値をより頻繁に再利用しますが、それも WW がそれを保持すべきキャッシュを超えるまでです。このトレードオフこそがチューニングの探索対象であり、workspace_bytes は WW で、valid_candidate が厳格な上限として使います。既定の候補で最大の 128×128×64128 \times 128 \times 64 には W=8(8192+8192+16384)=262144W = 8(8192 + 8192 + 16384) = 262144 バイトが必要です。11 K. Goto and R. A. van de Geijn, “Anatomy of high-performance matrix multiplication”, ACM TOMS 34(3), 2008.

ブロック化カーネルの厳密性

各要素について、gemm_reference は次を計算します

c^ij=fl(⋯fl(fl(0+ai0b0j)+ai1b1j)⋯+ai,k−1bk−1,j).\hat c_{ij} = \mathrm{fl}\bigl(\cdots\mathrm{fl}(\mathrm{fl}(0 + a_{i0} b_{0j}) + a_{i1} b_{1j}) \cdots + a_{i,k-1} b_{k-1,j}\bigr).

gemm_blocked は各要素を最初の kk ブロックで 0.00.0 から始め、部分和を CC に格納し(倍精度なので格納は厳密です)、次の kk ブロックで読み込み直し、ll の昇順に続けます。同じオペランドに対して同じ順序で同じ演算を行うため、コンパイラが両方のループを同様に扱う限り(特に、一方だけを融合積和演算に縮約しない限り)、2 つの結果はビット単位で同一です。パッキングは値をコピーするだけです。したがって、組み込みのカーネルの正しさの検査には許容誤差 00 を使えます。

総和の順序を変えるカーネル(ベクトル化されたマイクロカーネル、異なるループ順序)は、近い結果にしかなりません。単位丸め誤差 u=2−53u = 2^{-53} と γk=ku/(1−ku)\gamma_k = ku/(1 - ku) を用いると、任意の総和の順序について次が成り立ちます22 N. J. Higham, Accuracy and Stability of Numerical Algorithms, 2nd ed., SIAM, 2002, §3.1.

∣c^ij−cij∣≤γk∑l∣ail∣ ∣blj∣,\lvert \hat c_{ij} - c_{ij}\rvert \le \gamma_k \sum_l \lvert a_{il}\rvert\,\lvert b_{lj}\rvert ,

したがって、そのような 2 つのカーネルの差は高々 2γk∑l∣ail∣∣blj∣2\gamma_k \sum_l \lvert a_{il}\rvert\lvert b_{lj}\rvert です。reproducible_matrix の要素では ∣a∣,∣b∣≤1\lvert a\rvert, \lvert b\rvert \le 1 であり、上界は 2γkk≈2k2u2\gamma_k k \approx 2k^2 u です。k=64k = 64 ではおよそ 9.1⋅10−139.1\cdot 10^{-13} です。これが、これらの入力に対して順序を変えるカーネルに validate_gemm で渡すべき絶対許容誤差です。

再現可能な入力

reproducible_matrix(seed, r, c, layout) は 64 ビットの線形合同法生成器を実行します

xj+1=6364136223846793005 xj+1442695040888963407(mod264),x0=seed⊕0x9E3779B97F4A7C15,x_{j+1} = 6364136223846793005\, x_j + 1442695040888963407 \pmod{2^{64}}, \qquad x_0 = \text{seed} \oplus \mathtt{0x9E3779B97F4A7C15},

そして各状態を v=((x≫32) mod 2001)/1000−1v = \bigl((x \gg 32) \bmod 2001\bigr)/1000 - 1 に写します。乗数は ≡1(mod4)\equiv 1 \pmod 4 で増分は奇数なので、Hull–Dobell の定理により、生成器は最大周期 2642^{64} を持ちます。32 ビットの値 x≫32x \gg 32 は 2001 を法として簡約されます。232=2146410⋅2001+8862^{32} = 2146410 \cdot 2001 + 886 なので、2001 個の値のうち 886 個は相対的に 1/2146410≈4.7⋅10−71/2146410 \approx 4.7\cdot 10^{-7} だけわずかに出やすくなります。要素は [−1,1][-1, 1] にある 0.0010.001 の倍数です。論理的な行列は行優先の順序で生成されてから要求されたレイアウトで格納されるため、レイアウトによって値は変わりません。

端数

端のタイルは、ブロックサイズの倍数でない mm、nn、kk を処理します。boundary_shapes は m−1m - 1、mm、m+1m + 1 を返します。mm が mrmr の倍数なら、これらはそれぞれ短いタイル、ちょうど収まる場合、1 行の端数を試します。

設計上の決定

移植可能なスカラーカーネル

問題。 実際の GEMM チューニングは、ターゲットごとに異なる SIMD マイクロカーネルに依存します。選択。 1 つのスカラーのブロック化カーネルを使い、microkernel_id はラベルとします。理由。 このパッケージは、すべてのバックエンドでチューニングの仕組みを示しテストします。ハーネスはラベルを独自のカーネルに対応付けられます。

計測の前に検証する

gemm_is_correct は、候補が一度でも計時される前に候補を実行し、参照と比較します。execute_candidate は、無効な候補や形状に対して、高速で空の乗算に見えてしまう結果ではなく None を返します。

厳格な制約としてのワークスペース

候補が必要とするメモリはパラメータから計算され、実行前に上限と照合されるため、探索がターゲットのメモリ予算内で実行できない構成を選ぶことはありません。

明示的な計時範囲

TimingScope と AllocationMode は候補とその ID の一部なので、「計算のみ」と「エンドツーエンド」の結果が同じものを計測したかのように比較されることはありません。

シリアライズされた設定

config_json は、チューニングの実行を繰り返すのに必要なすべてを、64 ビット値は文字列として、スキーマバージョン mmkts_1 とともに書き出します。GemmEnvironment と合わせて、チューニングした結果がどこで有効かを示します。

正しさと不変条件

  • 上記の条件の下で gemm_blocked は gemm_reference とビット単位で等しくなります。テストは端数を含む 8 通りのレイアウトの組み合わせすべてを検査します。
  • すべての有効な候補について workspace_bytes(c) <= limit です。
  • enumerate_gemm_candidates は同じ 486 個の候補を同じ順序で返し、それらはすべて上限 262144 について有効です。
  • reproducible_matrix は引数だけに依存します。行優先と列優先の結果は同じ論理的な行列を表します。
  • validate_gemm は長さの差の要素をすべて不一致として数えるため、切り詰められた結果が有効になることはありません。

採用しなかった代替案

  • ターゲット固有のマイクロカーネル。 MoonBit のバックエンド間で移植できません。
  • 検証における相対許容誤差。 CC の要素はゼロに近いことがあります。有界な入力に対しては、上の絶対的な上界が誠実なものです。
  • プラットフォームの乱数生成器で入力を生成する。 ターゲット間で再現できません。

境界

  • ループ順序は m_n_k のみで、コードはスカラーのみです。3 つのマイクロカーネル ID は同じループを実行します。
  • pack_operand はオペランド全体を行優先の順序にコピーします。キャッシュサイズのパック済みパネルは構築しません。
  • allocation_mode、timing_scope、reuse_count は記述的なものです。カーネルは呼び出しのたびに結果を割り当て、何を計時するかはハーネスが決めます。
  • 探索ループはありません。探索を駆動するのは tune とアプリケーションです。
  • reproducible_matrix は 1 つの行列を作ります。バッチ化されたオペランドは呼び出し側が連結します。

Footnotes

  1. K. Goto and R. A. van de Geijn, “Anatomy of high-performance matrix multiplication”, ACM TOMS 34(3), 2008. ↩

  2. N. J. Higham, Accuracy and Stability of Numerical Algorithms, 2nd ed., SIAM, 2002, §3.1. ↩