淑樺說得好
你放不下的 別人已經進去了
(X
哈哈哈 今天還是ai slob 抱歉 昨天根莖葉
昨天跟今天的稿子還沒修好 麻煩先擔待一下
昨天我們看的是 single-block FWT。
如果 N <= 15,我們還可以很激進地把整個 chunk 放在一個 CUDA block 的 register/shared memory 範圍裡處理。這樣 global memory 只進出一次,中間大部分 butterfly 都在 register 或 warp shuffle 裡完成。
但 N 再往上就不行了。
N=16 的 vector 長度是:
2^16 = 65536 floats
FP32 大小是:
65536 * 4 bytes = 256 KB
這已經不可能靠一個 block 吃完整個 transform。
所以今天要切第二段:large-N multi-pass。
FWT 的 butterfly stage 有一個很麻煩的地方:stage 越後面,pair 距離越遠。
stage 0: 距離 1
stage 1: 距離 2
stage 2: 距離 4
...
stage s: 距離 2^s
小 stride 可以在 thread 或 warp 裡處理。
中等 stride 可以靠 shared memory。
但大 stride 會跨 block。
一旦跨 block,就不可能用 __syncthreads() 解決。CUDA block 之間沒有一般意義上的同步,也不能互相直接交換 register。
所以做 large-N FWT 時,真正的問題變成:
怎麼把原本一條長 vector 的 butterfly,
拆成幾個可以由 block 各自完成的 pass?
Hadamard 有一個很方便的結構:
H_N = H_a ⊗ H_b
如果 N = a + b,那長度 2^N 的 vector 可以看成一個矩陣:
rows = 2^a
cols = 2^b
接著整個 transform 可以拆成兩段:
先對每一段連續資料做 H_b
再沿著另一個維度做 H_a
直覺上就是:
Pass 1:row/local transform
Pass 2:column/inter-block transform
這不是為了把問題變漂亮,而是為了配合 GPU memory hierarchy。
Pass 1 盡量選一個 block 可以吃下的大小。
Pass 2 則處理剩下跨 chunk 的 butterfly。
如果第一段只吃 N=10,剩下 N=20 的 case 就會變成:
N_INTRA = 10
N_INTER = 10
Pass 2 要處理 10 個跨 chunk stage。這會很重。
如果第一段可以吃到 N=15,同樣 N=20 會變成:
N_INTRA = 15
N_INTER = 5
這差很多。
因為 N_INTER <= 5 時,剛好可以把 column 的 32 個 row 對應到一個 warp。
也就是說,Pass 2 可以用 warp shuffle 做:
ROWS = 2^N_INTER <= 32
這時候不需要 shared memory 做整個 column butterfly。
所以 N=15 single-block kernel 的意義不只是在 N=15 本身變快。
它也讓 N=16 到 N=20 的第二段變得很輕。
Pass 1 的資料 layout 是最友善的。
每個 chunk 是連續的:
chunk_size = 2^N_INTRA
所以 kernel 可以直接把每個 chunk 當成昨天的 single-block 問題。
大致是:
num_chunks = batch * state_len / chunk_size
for each chunk:
run native_fp32_extreme_fwt_kernel_v2
如果 N_INTRA = 15:
chunk_size = 32768 floats
每個 block 處理一個 chunk。
這段的好處是 global memory 讀寫是連續的,float4 load/store 很自然。
Pass 2 麻煩在 memory layout。
假設把 vector 看成:
rows = 2^N_INTER
cols = 2^N_INTRA
那 Pass 2 要對每個 column 做 FWT。
對同一個 column 來說,資料在 memory 裡的距離是:
cols
也就是 stride 很大。
這會讓 naive column access 很不連續。
所以 Pass 2 的 kernel 要小心兩件事:
1. 怎麼讓 global load/store 盡量 coalesced
2. 怎麼在 block 內重排成適合 butterfly 的形狀
當 N_INTER <= 5,column 高度最多是 32。
這剛好是一個 warp。
所以可以讓一個 warp 對一個 column 或一組 column 做:
float peer = __shfl_xor_sync(0xffffffff, val, stride);
每個 lane 拿一個 row 的值,stage 直接靠 XOR shuffle 交換。
這條路的優點很明確:
沒有 shared memory
沒有 bank conflict
butterfly exchange 在 warp 內完成
但它不是免費的。
你仍然要從 global memory 把 column 讀進來,也要寫回去。只是在 block 內計算時,已經把額外 shared memory 成本壓掉。
如果 N_INTER > 5,column 高度超過 32。
這時候單一 warp 放不下,就要用 tiled shared memory。
case_3 裡有一條 inter-block kernel:
native_fp32_interblock_fwt_v2_kernel<TILE_COLS, ...>
它的方向是:
一次處理 TILE_COLS 個 column
用 float4 做 global load
放進 shared memory tile
在 tile 裡做 row-direction butterfly
最後 float4 store 回去
這裡的 tuning 點不是 Tensor Core,而是 memory transaction:
TILE_COLS 太小 -> bandwidth utilization 不好
TILE_COLS 太大 -> shared memory 壓力上升
pitch 沒補好 -> bank conflict
load/store 不連續 -> sector efficiency 下降
這也是為什麼 large-N FWT 很難寫得漂亮。它不是算不動,而是資料不再自然連續。
case_3 的報告裡,N=15 和 N=16 的差距很明顯:
N=15: 82.98 us
N=16: 584.56 us
這不是因為 N=16 多一層加減法就多了這麼多時間。
真正原因是 N=16 開始需要 multi-pass。
以 batch=128 來估:
state size = 128 * 65536 * 4 bytes
~= 33.5 MB
如果做 2-pass,每個 pass 至少讀寫一次:
traffic ~= 33.5 MB * 4
~= 134 MB
報告裡 N=16 約 584 us,換算 effective bandwidth 大概是:
134 MB / 584 us ~= 229 GB/s
這代表 kernel 已經很接近在搬資料,而不是在等加減法。
所以 large-N 的優化方向不是「再塞更多 ALU」。
而是:
減少 pass 數
讓每次 pass 的 load/store 更連續
把 transpose/corner turning 變便宜
large-N 的驗證不能只驗小 N。
小 N 對了,只代表 single-block kernel 對。
large-N 還要驗:
Pass 1 output layout 是否正確
Pass 2 column mapping 是否正確
normalization 是在最後乘,還是每 pass 分開乘
batch stride 是否正確
in-place 讀寫有沒有覆蓋問題
case_3 曾經踩過一個很典型的問題:reinterpret_cast<float4*> 之後 pointer arithmetic 寫錯,造成 batch 間 memory overlap。
這種錯誤很危險,因為 kernel 可以跑,速度也可能看起來正常,但 output 已經被覆蓋。
所以驗收要有兩種:
小 N:對 naive PyTorch FWT
中大 N:對 reference FWT 或 norm conservation
norm conservation 只能當 sanity check,不能取代 elementwise diff。
large-N FWT 的核心不是新的數學,而是切分。
N <= 15:
single-block extreme kernel
N > 15:
Pass 1: intra-block FWT
Pass 2: inter-block / column FWT
能不能把 N_INTRA 推到 15,會直接影響後面 N_INTER 有多重。
如果 N_INTER <= 5,Pass 2 可以走 warp shuffle,這是非常好的情況。
如果超過,就只能回到 shared memory tiled kernel,開始跟 global memory traffic、coalescing、bank conflict 打架。
下一篇我們回頭看 Tensor Core:既然 Hadamard 只有 +1/-1,為什麼把它塞進 MMA 不一定會贏。
會贏的喔
新宿再次響起了五條悟的吟唱。“九綱”“偏光”“烏與聲明”“表裏之間”、宿儺明白自己再也沒有任何機會阻止茈的誕生了....
(X
明天見~