(笑死,前幾天都直接用不講,現在才突然間在那邊補這些知識,也不知道是什麼心態。但總之還是補了,雖然也沒有補得很好..
唉,這玩意兒是真的很難寫...AI slob寫什麼我發什麼吧 後續再 再 再看看...)
Day 10 我們把精度缺口補上了。
簡單說,前面那個可以跑的版本還不夠完整,因為:
模數太少
模數範圍太小
mantissa slice 不夠
pair schedule 不夠
如果要往 FP64 的 53-bit mantissa 靠近,就要把這些東西補起來。
但補起來以後,新的問題馬上出現:
精度變好了,工作量也變多了。
今天要講的就是這件事。
我們已經知道要做更多 modulus、更多 slice pair、更多 CRT reconstruction。那要怎麼讓它不要慢到失去意義?
這篇會先補一點 GPU 硬體知識,再回來看整個高精度 GEMM pipeline 要怎麼加速。
很多時候我們說「GPU 很快」,其實講得太粗。
真正寫 kernel 時,速度不是只看有沒有用到 GPU,而是看資料有沒有照 GPU 喜歡的方式流動。
一個簡化版的資料路徑大概是:
Device Global Memory
|
v
L2 Cache
|
v
Shared Memory / L1
|
v
Registers
|
v
Tensor Core / CUDA Core
|
v
Registers
|
v
Global Memory
每一層的特性不同。
| 位置 | 可以怎麼想 |
|---|---|
| Global Memory | 容量最大,但最慢,適合大量連續讀寫 |
| L2 Cache | 全 GPU 共用的 cache,幫 global memory 擋掉重複讀取 |
| Shared Memory | 每個 SM 裡的高速暫存空間,由 thread block 共用 |
| Register | 每個 thread 自己用,最快,但數量有限 |
| Tensor Core | 專門做小矩陣乘法的硬體單元 |
| CUDA Core / ALU | 做一般整數、浮點、位元運算 |
高效能 kernel 的目標通常不是「少寫幾行 code」。
目標是:
讓資料從 global memory 進來以後,在 shared memory / register 裡被重複使用,盡量少來回寫回 global memory。
這也是 FlashAttention 那種設計的核心精神。
不是因為它用了某個神奇指令,而是因為它避免把巨大的中間矩陣反覆寫回 HBM/global memory。
Tensor Core 可以先想成 GPU 裡的矩陣乘法專用硬體。
一般 CUDA core 比較像在做 scalar 運算:
sum += a * b;
Tensor Core 則是讓一個 warp 合作做一小塊矩陣乘法。
例如前面提到的 INT8 MMA 指令可以想成:
mma.sync.m16n8k32.s8
它的意思大概是:
A fragment: 16 x 32, int8
B fragment: 32 x 8, int8
C fragment: 16 x 8, int32 accumulator
也就是說,它不是一個 thread 算一個乘法。
它是整個 warp 一起餵資料給 Tensor Core,讓 Tensor Core 做一小塊矩陣乘法,結果累在 register 裡的 accumulator。
所以要用好 Tensor Core,重點不是只把資料型別換成 int8_t。
還要做這些事:
global memory 讀取要連續
shared memory layout 要對
ldmatrix 要讀到正確的 fragment
warp 裡每個 lane 要拿到正確資料
mma.sync 的 tile 要排得夠密
accumulator 不要被太多 epilogue 拖住
如果這些沒做好,程式碼裡有 mma.sync,Tensor Core 也不一定會忙。
cp.async、ldmatrix、mma.sync 各自在做什麼Day 9 已經碰過這三個名字。今天把它們放到硬體路徑裡看。
cp.asynccp.async 可以把 global memory 的資料搬到 shared memory。
重點是 async。
也就是說,kernel 可以一邊算目前 tile,一邊預先搬下一個 tile:
算 K tile 0
同時搬 K tile 1
算 K tile 1
同時搬 K tile 2
這叫 pipelining。
如果做得好,global memory latency 會被藏起來。
如果做不好,Tensor Core 就會等資料。
ldmatrixTensor Core 不直接吃普通 row-major 陣列。
它要吃的是 warp fragment。
ldmatrix 做的事就是讓 warp 從 shared memory 裡,用 Tensor Core 需要的方式把資料載進 register。
所以 shared memory layout 不是隨便排。
你要讓 ldmatrix 讀的時候:
lane mapping 正確
bank conflict 少
資料對齊
不然就算 global memory 搬得很快,shared memory 到 register 這段還是會卡。
mma.syncmma.sync 是真正呼叫 Tensor Core 做矩陣乘法的地方。
它吃 A fragment、B fragment、C accumulator,吐回新的 C accumulator。
在我們這種 INT8 GEMM 裡,最重要的是:
int8 x int8 -> int32 accumulator
後面 CRT 要用的就是這些 int32 accumulator。
Day 10 說過,如果要補齊精度,我們要做更多東西。
假設有:
P 個 modulus
Q 個 slice pair
那最直覺的 staged pipeline 會長這樣:
for each modulus:
產生 A residue
產生 B residue
for each slice pair:
INT8 GEMM
CRT / accumulate
這樣最好 debug,因為每一層都可以單獨驗證。
但它會慢在幾個地方:
A/B residue 寫回 global memory
C_mod 寫回 global memory
每個 modulus 或 pair 都有 kernel launch
CRT reconstruction 做太多 scalar work
register 用太多導致 occupancy 下降
所以正確版通常會經過兩個階段:
先拆開,讓答案對
再融合,讓資料少搬
這個順序很重要。
如果一開始就把全部塞進一顆巨大 kernel,debug 會非常痛苦。你不知道錯的是切片、模數、GEMM、CRT、scaling,還是 memory layout。
直覺上我們會想:
那就把所有東西都融合成一顆 kernel,不就最快?
不一定。
融合有兩種結果。
好的融合:
少寫 global memory
少讀中間結果
少 kernel launch
hot data 留在 register/shared memory
壞的融合:
register pressure 爆掉
occupancy 掉下去
scalar epilogue 卡住 Tensor Core
shared memory 太大,SM 同時跑不了幾個 block
在我們這種高精度 GEMM 裡,最危險的是 CRT epilogue。
Tensor Core 負責的是:
int8 x int8 -> int32
但 CRT 重建常常是:
mod
inverse
mixed radix
wide integer reconstruction
scaling back to double
這些不是 Tensor Core 做的。
它們多半落在 CUDA core / integer ALU / FP64 pipeline 上。
如果把太多 CRT 工作塞進同一顆 MMA kernel,結果可能會變成:
Tensor Core 很快算完
後面 scalar epilogue 做太久
整個 warp 被 epilogue 拖住
這種融合看起來很漂亮,但不一定快。
高精度版要變快,不能只喊 fused。
比較合理的方向是逐步把最貴的 global memory round-trip 拿掉。
最一開始應該保留清楚 stage:
scale / exponent
split / residue
INT8 GEMM
CRT
undo scaling
每一層都可以單獨測。
這時候慢是正常的。
它的價值是讓我們知道數學對不對。
如果每個 modulus 都單獨跑一輪 GEMM,launch 和 memory traffic 會很多。
比較好的方式是把幾個 modulus 合成一組:
modulus group 0: p0, p1, p2, p3
modulus group 1: p4, p5, p6, p7
...
這樣一顆 kernel 可以處理多個 modulus。
但 group size 不能無限加大,因為每多一個 modulus,就多一組 accumulator / residue / CRT 狀態。
這會吃 register。
如果每個 modulus 的完整 C_mod 都寫回 global memory,最後再一次讀回來做 CRT,memory traffic 會很大。
比較好的方向是:
一組 modulus 做完 MMA
就在 accumulator 附近先做 partial CRT
只把 group partial 寫回去
也就是:
r0, r1, r2, r3
-> partial CRT
-> x_group
這樣寫回 global memory 的資料量會變少,後面的 combine kernel 也比較便宜。
這個做法很像 FlashAttention 的精神:
不是把所有中間結果完整 materialize,而是每個 tile 做完就立刻消化掉,只留下後面真的需要的狀態。
A 和 B 的 residue 產生本來是兩段獨立工作。
如果硬體資源允許,可以讓它們用不同 stream 重疊:
stream 0: quantize A
stream 1: quantize B
這不是萬靈丹,因為兩邊還是會搶 memory bandwidth 和 SM。
但如果原本 A/B quantize 是明顯的 staged overhead,重疊可以回收一部分時間。
Tensor Core kernel 通常會有 CTA tile,例如:
128 x 128 x 64
128 x 64 x 64
64 x 128 x 64
tile 越大,資料重用可能越好。
但 tile 越大,shared memory 和 register 壓力也越大。
所以 tile tuning 不是看哪個形狀最帥,而是看:
Tensor pipe utilization
occupancy
registers per thread
shared memory per block
eligible warps
memory throughput
如果 tile 變小讓 occupancy 上升,但需要更多 CTA、更多 epilogue、更多 launch work,最後也可能更慢。
看起來最理想的做法是:
不要先把 A_mod/B_mod 寫到 global memory
每個 CTA 直接讀 double A/B tile
在 shared memory 或 register 裡轉成 residue
馬上餵給 MMA
這就是更 aggressive 的 FlashAttention-style 想法。
但它也有風險。
因為 double 到 residue 的轉換不是免費的:
讀 double
scale
round
mod
轉 balanced int8
如果每個 K tile 都重做這些 scalar 工作,省下的 memory traffic 可能抵不過新增的 quantize 成本。
所以 tile-local prologue 應該先當 prototype,不應該一開始就變成預設。
優化到這個階段,不能只看總時間。
總時間告訴你「有沒有變快」。
NCU 告訴你「為什麼」。
幾個很重要的指標:
| 指標 | 代表什麼 |
|---|---|
| Tensor pipe utilization | Tensor Core 忙不忙 |
| FP64 pipe utilization | FP64 單元是不是被 epilogue 卡住 |
| ALU pipe utilization | 整數/一般運算是不是太多 |
| DRAM throughput | 是否真的被 global memory bandwidth 卡住 |
| L2 hit rate | global memory 讀取是否有重用 |
| registers per thread | register pressure |
| shared memory per block | 每個 block 吃多少 shared memory |
| occupancy | 一個 SM 同時能跑多少 block/warp |
| eligible warps | scheduler 每 cycle 有多少 warp 可以發射 |
如果看到:
DRAM throughput 很高
Tensor pipe 很低
可能是 memory-bound。
如果看到:
DRAM throughput 很低
FP64 / ALU 很高
Tensor pipe 很低
那就不是 memory 搬太慢,而是 scalar epilogue 或 CRT 算太重。
這兩種瓶頸的解法完全不同。
memory-bound 要減少 global memory round-trip。
epilogue-bound 要減少 CRT/FP64 的 per-output 成本,或把它拆到更適合的階段。
Day 11 的重點是:精度補齊以後,速度問題會變成硬體資料流問題。
要真的變快,需要同時想數學和硬體。
Tensor Core 只負責 MMA,不負責整個 Ozaki/CRT。
mma.sync 可以加速 int8 x int8 -> int32,但 CRT、scaling、wide integer reconstruction 還是其他硬體單元在做。
資料搬運通常比你想的貴。
A/B residue、C_mod、partial result 如果一直寫回 global memory,Tensor Core 再快也會被拖住。
融合要有選擇。
把 partial CRT 靠近 accumulator 可能會快;把完整 CRT/FP64 epilogue 全塞進 MMA kernel 可能反而讓 Tensor Core 閒著。
tile tuning 要看 NCU。
tile 大小、stage 數、swizzle、group size 都要用 wall time 和 profiler 決定。
FlashAttention 精神不是複製 attention。
它真正值得學的是:不要 materialize 不必要的大中間矩陣,讓 tile 在 on-chip memory 裡被消化掉。
今天我們把「怎麼加速」的方向定清楚。
明天 Day 12 就可以公平地看數字:跟 native DGEMM、cuBLAS emulation、不同 fused 設計相比,到底輸在哪裡、贏在哪裡。