iT邦幫忙

2026 iThome 鐵人賽

DAY 16
0
Software Development

在 AI Compiler 工程師的路上系列 第 16

Day15:TTIR:Triton 前端如何理解矩陣乘法

  • 分享至 

  • xImage
  •  

昨天已經在 ASUS Ascent GX10 上產生一組固定 artifact,今天先從 TTIR 開始,看 Python kernel 進入 compiler 後,program id、二維 pointer、mask、K loop 與矩陣乘法還認不認得出來。

本文的 IR 片段取自這次實際編譯結果。為了讓版面可讀,我移除了重複的 loc(...) 與部分 attribute,並縮短少數 SSA 名稱,operation、type、shape、常數與資料流均保留。

完整看可以看 matmul_kernel.sourcematmul_kernel.ttir。前者是 AST 轉成 TTIR 後的初始 snapshot,後者是完成 TTIR passes 的版本。

本篇大綱

  • 分清 Python AST 產生初始 IR,以及 TTIR pass 最佳化這兩段。
  • 對照 .source.ttir,看每個 pass 實際最佳化了什麼。
  • 從 artifact 中找到 TTIR function。
  • 對照 tl.program_idtl.arange 與 pointer 計算。
  • 追 K loop、masked load 和 tt.dot
  • 觀察 specialization 如何折疊 compile-time constant。
  • 整理 TTIR 能回答的 correctness 問題。

先分清兩段轉換

Triton 3.6.0 的 Python 前端先走這條路

Python function AST
  → ASTSource.make_ir()
  → ast_to_ttir()
  → CodeGenerator.visit(fn.parse())
  → 初始 Triton dialect module(本實驗的 .source)

Triton 3.6.0 的 ASTSource.make_ir() 在 compiler.py 第 78~81 行呼叫 ast_to_ttir()ast_to_ttir() 的完整實作在 code_generator.py 第 1600~1639 行,其中第 1631 行generator.visit(fn.parse()) 開始走訪 kernel AST。

CodeGenerator 走訪 Python AST,呼叫 MLIR builder 建立 tt.functt.loadtt.dotscf.for 等 operation。這是前端 AST lowering,不是 MLIR pass。初始 module 建好後,NVIDIA backend 才呼叫 CUDABackend.make_ttir()

matmul_kernel.source
  → make_ttir() 中的 pass manager
  → matmul_kernel.ttir

Compiler 在 compiler.py 第 304~324 行先呼叫 src.make_ir(),把初始 module 保存為 .source,再進入後續 stages。

Triton 3.6.0 的 make_ttir() 實作在 NVIDIA backend compiler.py 第 230~245 行。直接列出 pass manager 的加入順序,可以看下的表對照。

順序 Triton 呼叫 這個 pass 負責什麼 本次 artifact 看到的結果
1 passes.common.add_inliner 展開可 Inlining 的 helper function cdivzerostt.call 消失
2 passes.ttir.add_rewrite_tensor_pointer 重寫 tensor-pointer operation 本 kernel 沒有 make_tensor_ptr 主線,無法從最終 diff 指認獨立變化
條件式 add_rewrite_tensor_descriptor_to_pointer 將 tensor descriptor 改寫為 pointer 條件是 capability // 10 < 9;SM121 不執行
3 passes.common.add_canonicalizer 折疊常數、消除冗餘 operation helper call 結果變成常數,部分 cast 與整數運算被簡化
4 passes.ttir.add_combine 合併 Triton dialect 中可簡化的計算 與 canonicalizer/CSE 共同縮短 pointer 與 mask 周邊計算
5 passes.ttir.add_reorder_broadcast 調整 broadcast 與其他 operation 的順序 最終 artifact 可見 broadcast 路徑較集中;單一 pass 效果需 per-pass dump 才能確定
6 passes.common.add_cse 合併重複子運算式 M/N 原本各有一個 0..63 range,最後共用一個
7 passes.common.add_symbol_dce 刪除不再被使用的 symbol Inlining 後的 private cdiv/zeros function 被移除
8 passes.ttir.add_loop_unroll 在合法且有利時展開迴圈 最終 TTIR 仍有 scf.for,本層沒有將 K loop 完全展開

補充一下這裡是多個 pass 會彼此接力,沒有保存每個單獨 pass 之後的結果。

.source → .ttir:具體變了什麼

Helper call 變成常數

.source 中,tl.cdiv(256, 64)tl.zeros((64, 64), tl.float32) 還是對 private function 的 tt.call。經過 inliner 與 canonicalization 後,.ttir 直接留下

%num_pid_n = arith.constant 4 : i32
%accumulator = arith.constant dense<0.000000e+00>
    : tensor<64x64xf32>

後續 symbol_dce 再把沒有 caller 的 private helper function 刪掉。因此這個差異不是一個 pass 單獨完成,而是「Inlining → 常數化簡 → 死 symbol 清除」。

兩個 range 變成一個

.source 為 M 與 N 各建立一次 tt.make_range {start = 0, end = 64}.ttir 只留下一個 %range_64,兩條 offset 路徑共用。這是 CSE 最容易用 artifact diff 驗證的例子。

整數與 pointer 鷹架被壓縮

.source 有多組 arith.extsi、整數 min/max 比較、overflow 處理與 ub.poison。到 .ttir 後,已特化的 shape/stride 讓許多條件可被證明,所以 pointer 與 loop bound 路徑明顯變短。這是 canonicalizer、combine 與 CSE 共同作用的結果;若要精確追到每一步,需要另存 per-pass IR dump。

還沒有變的部分也很重要

.ttir 仍保留 scf.fortt.loadtt.dot 與沒有 encoding 的 tensor type。這表示 TTIR pipeline 主要在清理與規範化 Triton 運算語意,尚未決定 warp layout、shared-memory staging 或具體 MMA instruction。這些變化是下一站 .ttir → .ttgir

第一次讀 TTIR,先看四條主線

TTIR 會出現大量 %a_ptrs_18%cst_3 這類 SSA 名稱。這些編號由 compiler 產生,第一次閱讀不用逐一記住。先追四條資料流

program id → M/N tile 座標
range → A/B/C pointer tensor
K loop → load → dot → accumulator
accumulator → cast → store

為什麼先看這四條?一個 matmul kernel 至少要回答四個 correctness 問題,這個 program instance 負責 C 的哪一塊、它會讀 A/B 的哪些位址、它如何沿 K 維累加,以及結果最後寫到 C 的哪裡。這四條資料流剛好可以回答這四個問題。

這個順序也能縮小除錯範圍,如果 tile 座標錯了,後面的 pointer 全會跟著錯,tile 座標正確但讀到錯誤元素,先查 shape、stride 與 pointer,讀取正確但數值仍錯,再查 K loop、mask、tt.dot 與 accumulator dtype;前三者都正確,最後才查 cast、store pointer 與 C mask。

從 function 入口開始

TTIR 的 function 開頭是

tt.func public @matmul_kernel(
    %a_ptr: !tt.ptr<f16> {tt.divisibility = 16 : i32},
    %b_ptr: !tt.ptr<f16> {tt.divisibility = 16 : i32},
    %c_ptr: !tt.ptr<f16> {tt.divisibility = 16 : i32},
    ...) {
  %pid = tt.get_program_id x : i32

三個主要 runtime argument 是 A、B、C 的 device pointer。M、N、K、stride 與 tile size 都是 tl.constexpr,這次 specialization 已經把它們折疊成常數,因此不會以一般 kernel parameter 的形式保留。

例如 TTIR 開頭直接出現

%num_pid_n = arith.constant 4 : i32
%c32_i32 = arith.constant 32 : i32
%c256_i32 = arith.constant 256 : i32
%c64_i32 = arith.constant 64 : i32
%accumulator = arith.constant dense<0.000000e+00>
    : tensor<64x64xf32>

這裡能驗證兩件事,矩陣大小與 tile size 已經 specialization,accumulator 是 64×64 的 FP32 tensor。若把 BLOCK_SIZE_K 改成 64,再編譯出來的 K range 和 loop step 也應跟著改變。

一個 program instance 對應哪塊 C

原始程式碼先把一維 pid 拆成二維 tile 座標

num_pid_n = tl.cdiv(n, block_size_n)  # 4
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n

TTIR 中可以找到對應的 signed division 與 remainder

%num_pid_n = arith.constant 4 : i32
%pid = tt.get_program_id x : i32
%pid_m = arith.divsi %pid, %num_pid_n : i32
%pid_n = arith.remsi %pid, %num_pid_n : i32

arith.divsi 是整數除法,arith.remsi 是取餘數。它們正好對應 Python 的 //%

接著建立長度 64 與 32 的 range。M、N 都是 64,因此 compiler 可以重用同一個 0..63 range

%range_64 = tt.make_range {start = 0 : i32, end = 64 : i32}
    : tensor<64xi32>
%offsets_k = tt.make_range {start = 0 : i32, end = 32 : i32}
    : tensor<32xi32>

tensor<64xi32> 它先表示一塊 tensor index。
如何分給 warp 與 lane,要到明天的 TTGIR 才決定。

二維 pointer 是怎麼形成的

Python 使用 [:, None][None, :] 組出二維 broadcast

a_ptrs = (
    a_ptr
    + offsets_m[:, None] * stride_am
    + offsets_k[None, :] * stride_ak
)

這裡的 a_ptrs 容易被誤讀成一個新配置的二維陣列。它是一張「由位址組成的 tensor」:a_ptrs[i, j] 存的是這個 program instance 應該讀取的 A 元素位址。之後的 tl.load(a_ptrs, ...) 會依這張位址表取回 64×32 個 FP16 值。

先把單一元素的位址寫出來會比較清楚。對 A 而言

a_ptrs[i, j]
  = a_ptr
  + offsets_m[i] * stride_am
  + offsets_k[j] * stride_ak

沿用昨天的 pid=6,這個 program instance 負責 M 方向的第 64..127 列。第一輪 K tile 的範圍是 0..31。如果 A 是一般 row-major 的 256×256 tensor,stride_am=256stride_ak=1,那麼

a_ptrs[2, 3]
  = a_ptr + 66 * 256 + 3
  = A[66, 3] 的位址

所以 i 選 A 的 row,j 選這一輪的 K 座標。把所有 i=0..63j=0..31 的組合列出來,就得到 64×32 個位址。

Python 的 offsets_moffsets_k 原本都是一維 tensor,兩者不能直接表達「每一 row 配上每一個 K 座標」。[:, None] 插入長度為 1 的 axis,讓 shape 變成 64×1[None, :] 則變成 1×32。Broadcast 再沿著長度為 1 的 axis 複製索引

offsets_m[:, None]       offsets_k[None, :]
shape = 64×1             shape = 1×32
      │ broadcast              │ broadcast
      ▼                        ▼
每一 row 重複 32 次          每個 K 座標重複 64 次
      └────────── 相加後得到 64×32 ──────────┘

TTIR 不會保留 Python indexing 語法,而會展開為

tt.expand_dims
→ tt.broadcast
→ arith.muli / arith.addi
→ tt.splat pointer
→ tt.addptr

以 A 為例,下面保留實際 TTIR 中會影響 pointer 資料流的步驟,只縮短 SSA 名稱與 location。%stride_am 是 shape 為 64×1、每個元素皆為 256 的常數 tensor

%a_rows = tt.expand_dims %offsets_m {axis = 1}
    : tensor<64xi32> -> tensor<64x1xi32>
%a_row_offsets = arith.muli %a_rows, %stride_am
    : tensor<64x1xi32>
%a_base = tt.splat %a_ptr
    : !tt.ptr<f16> -> tensor<64x1x!tt.ptr<f16>>
%a_row_ptrs = tt.addptr %a_base, %a_row_offsets
    : tensor<64x1x!tt.ptr<f16>>, tensor<64x1xi32>
%k_cols = tt.expand_dims %offsets_k {axis = 0}
    : tensor<32xi32> -> tensor<1x32xi32>
%a_row_ptr_grid = tt.broadcast %a_row_ptrs
    : tensor<64x1x!tt.ptr<f16>>
      -> tensor<64x32x!tt.ptr<f16>>
%k_offset_grid = tt.broadcast %k_cols
    : tensor<1x32xi32> -> tensor<64x32xi32>
%a_ptrs = tt.addptr %a_row_ptr_grid, %k_offset_grid
    : tensor<64x32x!tt.ptr<f16>>, tensor<64x32xi32>

在這組已特化的 row-major stride 中,stride_ak=1,所以最後加上的 %k_offset_grid 就是 K 座標本身。tt.splat 把單一 %a_ptr 複製成 pointer tensor,tt.addptr 則對其中每個 pointer 加上對應的元素 offset。這些 operation 沒有在這裡讀取 A,直到 tt.load 才會存取記憶體。

這段要從 type 讀

64×1 row index × stride_am=256
  → 64×1 row pointer
  + broadcast 後的 64×32 K offset
  → 64×32 A pointer tile

讀這一段時先比 shape

  • A pointer tile:tensor<64x32x!tt.ptr<f16>>
  • B pointer tile:tensor<32x64x!tt.ptr<f16>>
  • accumulator:tensor<64x64xf32>

三者剛好符合 (64×32) @ (32×64) → (64×64)。若 K 維放反、stride 寫錯或 broadcast axis 用錯,這裡通常比 SASS 更容易發現。

B 走同一套規則,只是兩個 axis 對調:b_ptrs[r, j] 指向 B[k_start+r, offsets_n[j]],所以 shape 是 32×64。編譯器需要保留這兩張 pointer tensor,才能讓後面的 tt.load 產生可交給 tt.dot 的兩個 tile。

K loop 與 masked load

這個例子 K=256BLOCK_SIZE_K=32,TTIR 保留結構化迴圈

%result:3 = scf.for %k_start = %c0_i32 to %c256_i32
    step %c32_i32
    iter_args(
      %a_ptrs_iter = %a_ptrs,
      %b_ptrs_iter = %b_ptrs,
      %acc_iter = %accumulator
    ) {

為什麼需要 K loop?C 的一個元素是沿 K 維做 reduction

C[m, n] = Σ A[m, k] × B[k, n],k = 0..255

這個 kernel 一次只把 K 維的 32 個元素做成 A/B tile,因此一輪只能為每個 C 元素累加 32 個乘積。256 / 32 = 8,需要八輪才能覆蓋完整 K 維

第 0 輪:k_start =   0,讀 K[  0: 32]
第 1 輪:k_start =  32,讀 K[ 32: 64]
...
第 7 輪:k_start = 224,讀 K[224:256]

scf.for 同時攜帶三個 loop-carried value,目前的 A pointer、B pointer,以及 accumulator。每輪結尾用 scf.yield 把更新後的三個值交給下一輪

這是 SSA 表示法處理變數更新的方式。IR 中的 SSA value 一旦定義就不會被原地修改,所以 Python 的 a_ptrs += ...accumulator = ... 會變成新的 %next_a_ptrs%next_acc,再由 scf.yield 傳給下一輪。每輪 A pointer 沿 K 維前進 32 * stride_ak,B pointer 則前進 32 * stride_bk

Loop body 先組出 mask,再執行 load。A 的真實資料流是

%m_ok = arith.cmpi slt, %offsets_m, %M
%k_ok = arith.cmpi slt, %k_start_plus_offsets, %K
%a_mask = arith.andi %m_ok_broadcast, %k_ok_broadcast
%a = tt.load %a_ptrs_iter, %a_mask, %zero
    : tensor<64x32x!tt.ptr<f16>>
%b = tt.load %b_ptrs_iter, %b_mask, %zero
    : tensor<32x64x!tt.ptr<f16>>

Mask 是一張與 pointer tile 對齊的 boolean tensor。以 A 的某個元素 [i, j] 為例,只有下面兩個條件同時成立才允許 load

offsets_m[i] < M
k_start + offsets_k[j] < K

第一個條件防止超出 A 的 row 數,第二個條件防止超出 A 的 K 維。B 則檢查 K 與 N。arith.andi 將兩個逐元素條件合成,因為任一座標越界,該位址就不能讀。

Mask 為 false 時,tt.load 不使用該越界位址取值,並回傳 other=0.0。零是乘加的中性補值,越界位置進入 tt.dot 後不會增加 accumulator。假設 K=250、步長仍是 32,最後一輪從 k_start=224 開始;j=0..25 對應有效的 K 座標 224..249j=26..31 會被 mask 掉並補零。這讓同一個固定 shape 的 64×32 tile 也能處理不整除 32 的尾端。

即使目前 256 能被 tile 整除,mask 的資料流仍來自原始程式。這能直接核對

  • A 的 mask 是否限制 M 與 K
  • B 的 mask 是否限制 K 與 N
  • other=0.0 是否成為越界 load 時填入的值

tl.dot 還沒有消失

迴圈最重要的一行是

%next_acc = tt.dot %a, %b, %acc_iter
    : tensor<64x32xf16> * tensor<32x64xf16>
      -> tensor<64x64xf32>
%next_a_ptrs = tt.addptr %a_ptrs_iter, %step_a
%next_b_ptrs = tt.addptr %b_ptrs_iter, %step_b
scf.yield %next_a_ptrs, %next_b_ptrs, %next_acc

tt.dot 在這裡代表一個 tile-level 的矩陣乘加。把它展開到單一輸出元素,語意是

next_acc[i, j]
  = acc_iter[i, j]
  + Σ a[i, r] * b[r, j],r = 0..31

每一輪會為 64×64 個 accumulator 元素,各加入 32 個乘積。八輪依序處理不同的 K 區段,最後才得到完整的 64×64 C tile。第三個 operand %acc_iter 很重要,每一輪的新部分乘積會累加至 %acc_iter,並產生傳給下一輪的 accumulator。

這裡保留了矩陣 shape、輸入 dtype 與 FP32 accumulator。TTIR 尚未決定要用哪一代 Tensor Core 指令,因此目前能確認它是可供後端 lowering 的 dot,還不能直接稱為 mma.sync、WGMMA 或 tcgen05

這也回答了 tl.dot 為什麼到 TTIR 還沒有消失。TTIR 的工作是保留 Triton 程式的 tile 語意,讓 compiler 還看得見 (64×32) @ (32×64) → (64×64)、FP16 輸入與 FP32 累加。此時 IR 還沒有 warp/lane layout,也沒有選定具體硬體指令,如果太早把 tt.dot 拆成 thread-level 指令,後續 pass 會失去完整的 shape 與 reduction 資訊,難以選擇資料 layout、shared-memory staging 和目標 GPU 的矩陣乘法路徑。

到了 TTGIR,compiler 會先替 operand 與 accumulator 加上 GPU layout,安排資料如何交給 warp 與 lane,更後面的 lowering 才把這個抽象 dot 逐步改成目標 GPU 能執行的 operation。對讀 TTIR 的人來說,tt.dot 仍在有助於檢查數學 shape 與 dtype,不必同時處理硬體分工細節。

迴圈結束後,accumulator 轉成 FP16,再寫回 C

%c = arith.truncf %result#2
    : tensor<64x64xf32> to tensor<64x64xf16>
%c_mask = arith.andi %m_ok_broadcast, %n_ok_broadcast
tt.store %c_ptrs, %c, %c_mask

這裡的 #2 表示取出 scf.for 第三個結果,也就是 accumulator。前兩個結果是迴圈結束時的 A/B pointer,後面不再使用。

Source location 讓 IR 能回到 Python

先直接讀這份matmul_kernel.source 底部的三個定義

#loc1 = loc("matmul.py":45:24)
#loc2 = loc("matmul.py":46:27)
#loc3 = loc("matmul.py":47:19)

#loc1#loc2#loc3 是 MLIR 的 location attribute alias。它們只是替較長的 location 取一個短名稱,避免 operation 每次都顯示完整的檔名、行號與字元位置。前面的 # 表示 attribute alias

以第一個為例,可以拆成

#loc1                         location 的短名稱
loc(...)                      這是一個 source location
"matmul.py"                   來源檔案
45                            第 45 行
24                            該行從 0 開始計數的第 24 個字元位置

Triton 3.6.0 的 CodeGenerator走訪 Python AST 時,會把 AST node 的 linenocol_offset 傳給 MLIR builder。lineno 從 1 開始,col_offset 從 0 開始。這三個位置可以對回以下 expression

pid = tl.program_id(axis=0)
# #loc1 的 45:24:`axis` 從字元位置 24 開始

num_pid_n = tl.cdiv(n, block_size_n)
# #loc2 的 46:27:`block_size_n` 從字元位置 27 開始

pid_m = pid // num_pid_n
# #loc3 的 47:19:右側的 `num_pid_n` 從字元位置 19 開始

Location 指向的是 compiler 當下走訪的 AST node,所以字元位置可能落在呼叫引數或右側 operand,不一定落在整行開頭,也不一定指向左側變數名稱。

接著再看這些短名稱如何接到 operation

%pid = tt.get_program_id x : i32 loc(#loc68)
%num_pid_n = tt.call ... loc(#loc69)
%pid_m = arith.divsi %pid, %num_pid_n : i32 loc(#loc70)

#loc68 = loc("pid"(#loc1))
#loc69 = loc("num_pid_n"(#loc2))
#loc70 = loc("pid_m"(#loc3))

%pid 為例,追蹤順序是

operation 尾端的 loc(#loc68)
  → #loc68 加上名稱 "pid"
  → #loc1
  → matmul.py:45:24
  → tl.program_id(axis=0) 附近

loc("pid"(#loc1)) 是 named location。它在檔案位置外再附上 pid 這個名稱,讓讀者知道該 IR value 對應哪個 Python 變數。也就是說,#loc1 負責記錄在哪裡,#loc68 再補上這個值叫什麼。

每個 operation 尾端的 source location 記錄它來自哪一段程式。MLIR 要求每個 operation 都帶 location,如果確實沒有來源,也必須明確使用 unknown location。

為什麼 IR 需要這份資訊?經過 lowering 與最佳化後,Python 變數名稱可能消失,SSA 名稱也可能換成另一個編號。一行 Python 還可能展開成 tt.expand_dimstt.broadcastarith.additt.addptr 等多個 operation,Location 提供一條來源線索,讓這些新的 IR operation 仍能指回產生它們的 Python expression。

TTIR 適合回答什麼

檢查項目 TTIR 位置 結果
specialization 是否正確 常數 256、64、32 M/N/K 與 tile size 已固定
program id 如何對應 tile divsi pid, 4remsi pid, 4 16 個 program 排成 4×4
A/B pointer shape 是否可相乘 64×3232×64 pointer tensor K 維方向一致
邊界 mask 是否正確 A 使用 M/K,B 使用 K/N,C 使用 M/N 三個 memory access 各自限制正確 axis
K loop 是否正確 0 to 256 step 32 共執行 8 個 K tile
dot dtype 是否正確 FP16 × FP16 → FP32 accumulator 寫回前才截成 FP16

這份 TTIR 已經實際回答,grid id 被拆成 4×4 tile、K loop 從 0 走到 256 而且步長為 32、A/B tile shape 可相乘、accumulator 使用 FP32、寫回前轉成 FP16。

Warp 如何分工、shared memory 是否 swizzle、最後用了幾個 register,不是從 TTIR 推測。明天繼續聊 TTGIR,這些 GPU mapping 才會第一次具體出現。

參考資料


上一篇
Day14 : Triton JIT 流程
下一篇
Day 16 : TTGIR 流程
系列文
在 AI Compiler 工程師的路上17
圖片
  熱門推薦
圖片
{{ item.channelVendor }} | {{ item.webinarstarted }} |
{{ formatDate(item.duration) }}
直播中

尚未有邦友留言

立即登入留言