iT邦幫忙

2026 iThome 鐵人賽

DAY 17
0
Software Development

GPU 效能優化實戰:30 天從 Kernel 到 Profiling (重賽版)系列 第 17

Day 17:一個 block 放不下之後,FWT 要怎麼切 [重賽版]

  • 分享至 

  • xImage
  •  

淑樺說得好
你放不下的 別人已經進去了
(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。


不能 single-block 之後,問題變成資料重排

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?

Kronecker split:把一條 vector 看成矩陣

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=15

如果第一段只吃 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=16N=20 的第二段變得很輕。


Pass 1:連續 chunk 的 FWT

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:跨 chunk 的 column FWT

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:直接用 warp shuffle

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:還是要 tiled 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 很難寫得漂亮。它不是算不動,而是資料不再自然連續。


為什麼 N=16 會突然變慢

case_3 的報告裡,N=15N=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=16584 us,換算 effective bandwidth 大概是:

134 MB / 584 us ~= 229 GB/s

這代表 kernel 已經很接近在搬資料,而不是在等加減法。

所以 large-N 的優化方向不是「再塞更多 ALU」。

而是:

減少 pass 數
讓每次 pass 的 load/store 更連續
把 transpose/corner turning 變便宜

正確性驗證也要分 pass

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。


Day 17 的結論

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
明天見~


上一篇
# Day 16:先把 Hadamard 塞進一個 block ( (重賽版))
系列文
GPU 效能優化實戰:30 天從 Kernel 到 Profiling (重賽版)17
圖片
  熱門推薦
圖片
{{ item.channelVendor }} | {{ item.webinarstarted }} |
{{ formatDate(item.duration) }}
直播中

尚未有邦友留言

立即登入留言