inference/kernel.py
21.7 KB · 537 lines · python Raw
1 import torch
2 import tilelang
3 import tilelang.language as T
4 from typing import Tuple, Optional
5
6
7 tilelang.set_log_level("WARNING")
8
9 pass_configs = {
10 tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True,
11 tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True,
12 }
13
14 FP8 = "float8_e4m3"
15 FP4 = "float4_e2m1fn"
16 FE8M0 = "float8_e8m0fnu"
17 BF16 = "bfloat16"
18 FP32 = "float32"
19 INT32 = "int32"
20
21
22 def fast_log2_ceil(x):
23 """Compute ceil(log2(x)) via IEEE 754 bit manipulation. Avoids slow log/ceil intrinsics."""
24 bits_x = T.reinterpret("uint32", x)
25 exp_x = (bits_x >> 23) & 0xFF
26 man_bits = bits_x & ((1 << 23) - 1)
27 return T.Cast("int32", exp_x - 127 + T.if_then_else(man_bits != 0, 1, 0))
28
29
30 def fast_pow2(x):
31 """Compute 2^x for integer x via IEEE 754 bit manipulation."""
32 bits_x = (x + 127) << 23
33 return T.reinterpret("float32", bits_x)
34
35
36 def fast_round_scale(amax, fp8_max_inv):
37 return fast_pow2(fast_log2_ceil(amax * fp8_max_inv))
38
39
40 @tilelang.jit(pass_configs=pass_configs)
41 def act_quant_kernel(
42 N, block_size=128, in_dtype=BF16, out_dtype=FP8, scale_dtype=FP32,
43 round_scale=False, inplace=False
44 ):
45 """Block-wise FP8 quantization. inplace=True does fused quant+dequant back to BF16."""
46 M = T.symbolic("M")
47 fp8_min = -448.0
48 fp8_max = 448.0
49 fp8_max_inv = 1 / fp8_max
50 num_stages = 0 if round_scale or inplace else 2
51 blk_m = 32
52 group_size = block_size
53 # Internal computation in FP32; scale_dtype controls output storage format.
54 compute_dtype = FP32
55 out_dtype = in_dtype if inplace else out_dtype
56
57 @T.prim_func
58 def act_quant_kernel_(
59 X: T.Tensor[(M, N), in_dtype],
60 Y: T.Tensor[(M, N), out_dtype],
61 S: T.Tensor[(M, T.ceildiv(N, group_size)), scale_dtype],
62 ):
63 with T.Kernel(T.ceildiv(M, blk_m), T.ceildiv(N, group_size), threads=128) as (
64 pid_m,
65 pid_n,
66 ):
67 x_shared = T.alloc_shared((blk_m, group_size), in_dtype)
68 x_local = T.alloc_fragment((blk_m, group_size), in_dtype)
69 amax_local = T.alloc_fragment((blk_m,), compute_dtype)
70 s_local = T.alloc_fragment((blk_m,), compute_dtype)
71 y_local = T.alloc_fragment((blk_m, group_size), out_dtype)
72 y_shared = T.alloc_shared((blk_m, group_size), out_dtype)
73
74 for _ in T.Pipelined(1, num_stages=num_stages):
75 T.copy(X[pid_m * blk_m, pid_n * group_size], x_shared)
76 T.copy(x_shared, x_local)
77 T.reduce_absmax(x_local, amax_local, dim=1)
78 for i in T.Parallel(blk_m):
79 amax_local[i] = T.max(amax_local[i], 1e-4)
80 if round_scale:
81 s_local[i] = fast_round_scale(amax_local[i], fp8_max_inv)
82 else:
83 s_local[i] = amax_local[i] * fp8_max_inv
84 if inplace:
85 for i, j in T.Parallel(blk_m, group_size):
86 y_local[i, j] = T.Cast(
87 out_dtype,
88 T.Cast(compute_dtype, T.Cast(FP8, T.clamp(
89 x_local[i, j] / s_local[i], fp8_min, fp8_max
90 ))) * s_local[i],
91 )
92 else:
93 for i, j in T.Parallel(blk_m, group_size):
94 y_local[i, j] = T.clamp(
95 x_local[i, j] / s_local[i], fp8_min, fp8_max
96 )
97 for i in T.Parallel(blk_m):
98 S[pid_m * blk_m + i, pid_n] = T.Cast(scale_dtype, s_local[i])
99 T.copy(y_local, y_shared)
100 T.copy(y_shared, Y[pid_m * blk_m, pid_n * group_size])
101
102 return act_quant_kernel_
103
104
105 def act_quant(
106 x: torch.Tensor, block_size: int = 128, scale_fmt: Optional[str] = None,
107 scale_dtype: torch.dtype = torch.float32, inplace: bool = False,
108 ) -> torch.Tensor:
109 """Block-wise FP8 quantization. inplace=True does fused quant+dequant back to BF16.
110 When scale_fmt is set, scales are rounded to power-of-2 (MXFP)."""
111 N = x.size(-1)
112 assert N % block_size == 0
113 tl_dtype = FE8M0 if scale_dtype == torch.float8_e8m0fnu else FP32
114 z = x.contiguous()
115 y = torch.empty_like(z) if inplace else torch.empty_like(z, dtype=torch.float8_e4m3fn)
116 s = z.new_empty(*z.size()[:-1], N // block_size, dtype=scale_dtype)
117 kernel = act_quant_kernel(
118 N, block_size, scale_dtype=tl_dtype,
119 round_scale=scale_fmt is not None, inplace=inplace,
120 )
121 kernel(z.view(-1, N), y.view(-1, N), s.view(-1, N // block_size))
122 if inplace:
123 x.copy_(y)
124 return x
125 return y, s
126
127
128 @tilelang.jit(pass_configs=pass_configs)
129 def fp4_quant_kernel(
130 N, block_size=32, in_dtype=BF16, scale_dtype=FE8M0, inplace=False
131 ):
132 """Block-wise FP4 quantization. Power-of-2 scale via bit ops. inplace=True does fused quant+dequant."""
133 M = T.symbolic("M")
134 fp4_max = 6.0
135 fp4_max_inv = 1.0 / fp4_max
136 blk_m = 32
137 group_size = block_size
138 compute_dtype = FP32
139 out_dtype = in_dtype if inplace else FP4
140
141 @T.prim_func
142 def fp4_quant_kernel_(
143 X: T.Tensor[(M, N), in_dtype],
144 Y: T.Tensor[(M, N), out_dtype],
145 S: T.Tensor[(M, T.ceildiv(N, group_size)), scale_dtype],
146 ):
147 with T.Kernel(T.ceildiv(M, blk_m), T.ceildiv(N, group_size), threads=128) as (
148 pid_m,
149 pid_n,
150 ):
151 x_shared = T.alloc_shared((blk_m, group_size), in_dtype)
152 x_local = T.alloc_fragment((blk_m, group_size), in_dtype)
153 amax_local = T.alloc_fragment((blk_m,), compute_dtype)
154 s_local = T.alloc_fragment((blk_m,), compute_dtype)
155 y_local = T.alloc_fragment((blk_m, group_size), out_dtype)
156 y_shared = T.alloc_shared((blk_m, group_size), out_dtype)
157
158 for _ in T.Pipelined(1, num_stages=2):
159 T.copy(X[pid_m * blk_m, pid_n * group_size], x_shared)
160 T.copy(x_shared, x_local)
161 T.reduce_absmax(x_local, amax_local, dim=1)
162 for i in T.Parallel(blk_m):
163 amax_local[i] = T.max(amax_local[i], 6 * (2**-126))
164 s_local[i] = fast_round_scale(amax_local[i], fp4_max_inv)
165 if inplace:
166 for i, j in T.Parallel(blk_m, group_size):
167 y_local[i, j] = T.Cast(
168 out_dtype,
169 T.Cast(compute_dtype, T.Cast(FP4, T.clamp(
170 x_local[i, j] / s_local[i], -fp4_max, fp4_max
171 ))) * s_local[i],
172 )
173 else:
174 for i, j in T.Parallel(blk_m, group_size):
175 y_local[i, j] = T.clamp(
176 x_local[i, j] / s_local[i], -fp4_max, fp4_max
177 )
178 for i in T.Parallel(blk_m):
179 S[pid_m * blk_m + i, pid_n] = T.Cast(scale_dtype, s_local[i])
180 T.copy(y_local, y_shared)
181 T.copy(y_shared, Y[pid_m * blk_m, pid_n * group_size])
182
183 return fp4_quant_kernel_
184
185
186 def fp4_act_quant(
187 x: torch.Tensor, block_size: int = 32, inplace: bool = False,
188 ) -> torch.Tensor:
189 """Block-wise FP4 quantization. inplace=True does fused quant+dequant back to BF16."""
190 N = x.size(-1)
191 assert N % block_size == 0
192 z = x.contiguous()
193 y = torch.empty_like(z) if inplace else z.new_empty(*z.shape[:-1], N // 2, dtype=torch.float4_e2m1fn_x2)
194 s = z.new_empty(*z.size()[:-1], N // block_size, dtype=torch.float8_e8m0fnu)
195 kernel = fp4_quant_kernel(N, block_size, inplace=inplace)
196 kernel(z.view(-1, N), y.view(-1, y.size(-1)), s.view(-1, N // block_size))
197 if inplace:
198 x.copy_(y)
199 return x
200 return y, s
201
202
203 @tilelang.jit(pass_configs=pass_configs)
204 def fp8_gemm_kernel(N, K, out_dtype=BF16, accum_dtype=FP32, scale_dtype=FP32):
205 assert out_dtype in [BF16, FP32]
206
207 M = T.symbolic("M")
208 group_size = 128
209 block_M = 32
210 block_N = 128
211 block_K = 128
212
213 @T.prim_func
214 def fp8_gemm_kernel_(
215 A: T.Tensor[(M, K), FP8],
216 B: T.Tensor[(N, K), FP8],
217 C: T.Tensor[(M, N), out_dtype],
218 scales_a: T.Tensor[(M, T.ceildiv(K, group_size)), scale_dtype],
219 scales_b: T.Tensor[(T.ceildiv(N, group_size), T.ceildiv(K, group_size)), scale_dtype],
220 ):
221 with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (
222 bx,
223 by,
224 ):
225 A_shared = T.alloc_shared((block_M, block_K), FP8)
226 B_shared = T.alloc_shared((block_N, block_K), FP8)
227 C_shared = T.alloc_shared((block_M, block_N), out_dtype)
228 Scale_C_shared = T.alloc_shared((block_M), FP32)
229 C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
230 C_local_accum = T.alloc_fragment((block_M, block_N), accum_dtype)
231
232 # Improve L2 Cache
233 T.use_swizzle(panel_size=10)
234 T.clear(C_local)
235 T.clear(C_local_accum)
236
237 K_iters = T.ceildiv(K, block_K)
238 for k in T.Pipelined(K_iters, num_stages=4):
239 T.copy(A[by * block_M, k * block_K], A_shared)
240 T.copy(B[bx * block_N, k * block_K], B_shared)
241 # Cast scales to FP32 for computation; scales_b has one value per block_N group
242 Scale_B = T.Cast(FP32, scales_b[bx * block_N // group_size, k])
243 for i in T.Parallel(block_M):
244 Scale_C_shared[i] = T.Cast(FP32, scales_a[by * block_M + i, k]) * Scale_B
245
246 T.gemm(A_shared, B_shared, C_local, transpose_B=True)
247 # Separate accumulator for scale-corrected results (2x accumulation precision)
248 for i, j in T.Parallel(block_M, block_N):
249 C_local_accum[i, j] += C_local[i, j] * Scale_C_shared[i]
250 T.clear(C_local)
251 T.copy(C_local_accum, C_shared)
252 T.copy(C_shared, C[by * block_M, bx * block_N])
253
254 return fp8_gemm_kernel_
255
256
257 def fp8_gemm(
258 a: torch.Tensor, a_s: torch.Tensor, b: torch.Tensor, b_s: torch.Tensor,
259 scale_dtype: torch.dtype = torch.float32,
260 ) -> torch.Tensor:
261 """C[M,N] = A[M,K] @ B[N,K]^T with per-128 block FP8 scaling on both A and B."""
262 assert a.is_contiguous() and b.is_contiguous(), "Input tensors must be contiguous"
263 assert a_s.is_contiguous() and b_s.is_contiguous(), (
264 "Scaling factor tensors must be contiguous"
265 )
266 tl_dtype = FE8M0 if scale_dtype == torch.float8_e8m0fnu else FP32
267 K = a.size(-1)
268 M = a.numel() // K
269 N = b.size(0)
270 c = a.new_empty(*a.size()[:-1], N, dtype=torch.get_default_dtype())
271 kernel = fp8_gemm_kernel(N, K, scale_dtype=tl_dtype)
272 kernel(a.view(M, K), b, c.view(M, N), a_s.view(M, -1), b_s)
273 return c
274
275
276 @tilelang.jit(pass_configs=pass_configs)
277 def sparse_attn_kernel(h: int, d: int, scale=None):
278 """Sparse multi-head attention via index gathering + online softmax (FlashAttention-style).
279 For each (batch, seq_pos), gathers top-k KV positions by index, computes attention
280 with numerically stable running max/sum, and includes a learnable attn_sink bias."""
281 b = T.symbolic("b")
282 m = T.symbolic("m")
283 n = T.symbolic("n")
284 topk = T.symbolic("topk")
285 if scale is None:
286 scale = (1.0 / d) ** 0.5
287
288 num_stages = 2
289 threads = 256
290 block = 64
291 num_blocks = tilelang.cdiv(topk, block)
292
293 @T.prim_func
294 def sparse_attn_kernel_(
295 q: T.Tensor[(b, m, h, d), BF16],
296 kv: T.Tensor[(b, n, d), BF16],
297 o: T.Tensor[(b, m, h, d), BF16],
298 attn_sink: T.Tensor[(h,), FP32],
299 topk_idxs: T.Tensor[(b, m, topk), INT32],
300 ):
301 with T.Kernel(m, b, threads=threads) as (bx, by):
302 q_shared = T.alloc_shared((h, d), BF16)
303 kv_shared = T.alloc_shared((block, d), BF16)
304 o_shared = T.alloc_shared((h, d), BF16)
305 acc_s_cast = T.alloc_shared((h, block), BF16)
306
307 idxs = T.alloc_fragment(block, INT32)
308 acc_s = T.alloc_fragment((h, block), FP32)
309 acc_o = T.alloc_fragment((h, d), FP32)
310 scores_max = T.alloc_fragment(h, FP32)
311 scores_max_prev = T.alloc_fragment(h, FP32)
312 scores_scale = T.alloc_fragment(h, FP32)
313 scores_sum = T.alloc_fragment(h, FP32)
314 sum_exp = T.alloc_fragment(h, FP32)
315
316 T.clear(acc_o)
317 T.clear(sum_exp)
318 T.fill(scores_max, -T.infinity(FP32))
319 T.copy(q[by, bx, :, :], q_shared)
320
321 for t in T.Pipelined(num_blocks, num_stages=num_stages):
322 for i in T.Parallel(block):
323 idxs[i] = T.if_then_else(t * block + i < topk, topk_idxs[by, bx, t * block + i], -1)
324 for i, j in T.Parallel(block, d):
325 kv_shared[i, j] = T.if_then_else(idxs[i] != -1, kv[by, idxs[i], j], 0)
326 for i, j in T.Parallel(h, block):
327 acc_s[i, j] = T.if_then_else(idxs[j] != -1, 0, -T.infinity(FP32))
328 T.gemm(q_shared, kv_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullRow)
329 for i, j in T.Parallel(h, block):
330 acc_s[i, j] *= scale
331 T.copy(scores_max, scores_max_prev)
332 T.reduce_max(acc_s, scores_max, dim=1, clear=False)
333 for i in T.Parallel(h):
334 scores_scale[i] = T.exp(scores_max_prev[i] - scores_max[i])
335 for i, j in T.Parallel(h, block):
336 acc_s[i, j] = T.exp(acc_s[i, j] - scores_max[i])
337 T.reduce_sum(acc_s, scores_sum, dim=1)
338 for i in T.Parallel(h):
339 sum_exp[i] = sum_exp[i] * scores_scale[i] + scores_sum[i]
340 T.copy(acc_s, acc_s_cast)
341 for i, j in T.Parallel(h, d):
342 acc_o[i, j] *= scores_scale[i]
343 T.gemm(acc_s_cast, kv_shared, acc_o, policy=T.GemmWarpPolicy.FullRow)
344
345 for i in T.Parallel(h):
346 sum_exp[i] += T.exp(attn_sink[i] - scores_max[i])
347 for i, j in T.Parallel(h, d):
348 acc_o[i, j] /= sum_exp[i]
349 T.copy(acc_o, o_shared)
350 T.copy(o_shared, o[by, bx, :, :])
351
352 return sparse_attn_kernel_
353
354
355 def sparse_attn(
356 q: torch.Tensor, kv: torch.Tensor, attn_sink: torch.Tensor, topk_idxs: torch.Tensor, softmax_scale: float
357 ) -> torch.Tensor:
358 b, s, h, d = q.size()
359 # Pad heads to 16 for kernel efficiency (stripped after)
360 if h < 16:
361 q = torch.cat([q, q.new_zeros(b, s, 16 - h, d)], dim=2)
362 attn_sink = torch.cat([attn_sink, attn_sink.new_zeros(16 - h)])
363 o = torch.empty_like(q)
364 kernel = sparse_attn_kernel(q.size(2), d, softmax_scale)
365 kernel(q, kv, o, attn_sink, topk_idxs)
366 if h < 16:
367 o = o.narrow(2, 0, h).contiguous()
368 return o
369
370
371 @tilelang.jit(pass_configs=pass_configs)
372 def hc_split_sinkhorn_kernel(hc: int, sinkhorn_iters: int, eps: float):
373 n = T.symbolic("n")
374 mix_hc = (2 + hc) * hc
375 threads = 64
376
377 @T.prim_func
378 def hc_split_sinkhorn_kernel_(
379 mixes: T.Tensor[(n, mix_hc), FP32],
380 hc_scale: T.Tensor[(3,), FP32],
381 hc_base: T.Tensor[(mix_hc,), FP32],
382 pre: T.Tensor[(n, hc), FP32],
383 post: T.Tensor[(n, hc), FP32],
384 comb: T.Tensor[(n, hc, hc), FP32],
385 ):
386 with T.Kernel(n, threads=threads) as i:
387 mixes_shared = T.alloc_shared(mix_hc, FP32)
388 comb_frag = T.alloc_fragment((hc, hc), FP32)
389 T.copy(mixes[i, :], mixes_shared)
390
391 for j in T.Parallel(hc):
392 pre[i, j] = T.sigmoid(mixes_shared[j] * hc_scale[0] + hc_base[j]) + eps
393 for j in T.Parallel(hc):
394 post[i, j] = 2 * T.sigmoid(mixes_shared[j + hc] * hc_scale[1] + hc_base[j + hc])
395 for j, k in T.Parallel(hc, hc):
396 comb_frag[j, k] = mixes_shared[j * hc + k + hc * 2] * hc_scale[2] + hc_base[j * hc + k + hc * 2]
397
398 row_sum = T.alloc_fragment(hc, FP32)
399 col_sum = T.alloc_fragment(hc, FP32)
400
401 # comb = comb.softmax(-1) + eps
402 row_max = T.alloc_fragment(hc, FP32)
403 T.reduce_max(comb_frag, row_max, dim=1)
404 for j, k in T.Parallel(hc, hc):
405 comb_frag[j, k] = T.exp(comb_frag[j, k] - row_max[j])
406 T.reduce_sum(comb_frag, row_sum, dim=1)
407 for j, k in T.Parallel(hc, hc):
408 comb_frag[j, k] = comb_frag[j, k] / row_sum[j] + eps
409
410 # comb = comb / (comb.sum(-2) + eps)
411 T.reduce_sum(comb_frag, col_sum, dim=0)
412 for j, k in T.Parallel(hc, hc):
413 comb_frag[j, k] = comb_frag[j, k] / (col_sum[k] + eps)
414
415 for _ in T.serial(sinkhorn_iters - 1):
416 # comb = comb / (comb.sum(-1) + eps)
417 T.reduce_sum(comb_frag, row_sum, dim=1)
418 for j, k in T.Parallel(hc, hc):
419 comb_frag[j, k] = comb_frag[j, k] / (row_sum[j] + eps)
420 # comb = comb / (comb.sum(-2) + eps)
421 T.reduce_sum(comb_frag, col_sum, dim=0)
422 for j, k in T.Parallel(hc, hc):
423 comb_frag[j, k] = comb_frag[j, k] / (col_sum[k] + eps)
424
425 T.copy(comb_frag, comb[i, :, :])
426
427 return hc_split_sinkhorn_kernel_
428
429
430 def hc_split_sinkhorn(mixes: torch.Tensor, hc_scale: torch.Tensor, hc_base: torch.Tensor, hc_mult: int = 4, sinkhorn_iters: int = 20, eps: float = 1e-6):
431 b, s, _ = mixes.size()
432 pre = mixes.new_empty(b, s, hc_mult)
433 post = mixes.new_empty(b, s, hc_mult)
434 comb = mixes.new_empty(b, s, hc_mult, hc_mult)
435 kernel = hc_split_sinkhorn_kernel(hc_mult, sinkhorn_iters, eps)
436 kernel(mixes.view(-1, (2 + hc_mult) * hc_mult), hc_scale, hc_base,
437 pre.view(-1, hc_mult), post.view(-1, hc_mult), comb.view(-1, hc_mult, hc_mult))
438 return pre, post, comb
439
440
441 @tilelang.jit(pass_configs=pass_configs)
442 def fp4_gemm_kernel(N, K, out_dtype=BF16, accum_dtype=FP32, scale_dtype=FP32):
443 """FP8 act x FP4 weight GEMM kernel.
444
445 C[M, N] = A_fp8[M, K] @ B_fp4[N, K]^T
446
447 Act: 1x128 quant on K (reduce dim), FP8 with configurable scale dtype
448 Weight: 1x32 quant on K (reduce dim), FP4 with E8M0 scale
449
450 B is stored as [N, K//2] in float4_e2m1fn_x2, logical [N, K] in fp4.
451 The FP4 values are packed along the K (last) dimension.
452
453 Strategy: load FP4 sub-blocks of size [block_N, sub_K] (sub_K=32),
454 cast FP4 to FP8 via float, then do FP8xFP8 GEMM.
455 Apply act scale (per 128 on K) and weight scale (per 32 on K) to the accumulator.
456 """
457 M = T.symbolic("M")
458 act_group_size = 128
459 weight_group_size = 32
460 block_M = 32
461 block_N = 128
462 block_K = 32 # matches weight_group_size for simple scale handling
463 n_sub = act_group_size // block_K # 4 sub-blocks per act scale group
464
465 @T.prim_func
466 def fp4_gemm_kernel_(
467 A: T.Tensor[(M, K), FP8],
468 B: T.Tensor[(N, K), FP4],
469 C: T.Tensor[(M, N), out_dtype],
470 scales_a: T.Tensor[(M, T.ceildiv(K, act_group_size)), scale_dtype],
471 scales_b: T.Tensor[(N, T.ceildiv(K, weight_group_size)), scale_dtype],
472 ):
473 with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (
474 bx,
475 by,
476 ):
477 A_shared = T.alloc_shared((block_M, block_K), FP8)
478 B_fp4_shared = T.alloc_shared((block_N, block_K), FP4)
479 B_shared = T.alloc_shared((block_N, block_K), FP8)
480 C_shared = T.alloc_shared((block_M, block_N), out_dtype)
481 C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
482 C_local_accum = T.alloc_fragment((block_M, block_N), accum_dtype)
483 scale_a_frag = T.alloc_fragment((block_M,), FP32)
484 scale_b_frag = T.alloc_fragment((block_N,), FP32)
485
486 T.use_swizzle(panel_size=10)
487 T.clear(C_local)
488 T.clear(C_local_accum)
489
490 K_iters = T.ceildiv(K, block_K)
491 for k in T.Pipelined(K_iters, num_stages=2):
492 T.copy(A[by * block_M, k * block_K], A_shared)
493 T.copy(B[bx * block_N, k * block_K], B_fp4_shared)
494 # FP4->FP8 cast must go through FP32 to avoid ambiguous C++ overload
495 for i, j in T.Parallel(block_N, block_K):
496 B_shared[i, j] = T.Cast(FP8, T.Cast(FP32, B_fp4_shared[i, j]))
497
498 # Weight scale: per 32 on K, indexed by k (each k is one block_K=32)
499 for i in T.Parallel(block_N):
500 scale_b_frag[i] = T.Cast(FP32, scales_b[bx * block_N + i, k])
501
502 # Act scale: per 128 on K, indexed by k // 4
503 for i in T.Parallel(block_M):
504 scale_a_frag[i] = T.Cast(FP32, scales_a[by * block_M + i, k // n_sub])
505
506 T.gemm(A_shared, B_shared, C_local, transpose_B=True)
507
508 for i, j in T.Parallel(block_M, block_N):
509 C_local_accum[i, j] += C_local[i, j] * scale_a_frag[i] * scale_b_frag[j]
510 T.clear(C_local)
511
512 T.copy(C_local_accum, C_shared)
513 T.copy(C_shared, C[by * block_M, bx * block_N])
514
515 return fp4_gemm_kernel_
516
517
518 def fp4_gemm(
519 a: torch.Tensor, a_s: torch.Tensor, b: torch.Tensor, b_s: torch.Tensor,
520 scale_dtype: torch.dtype = torch.float32,
521 ) -> torch.Tensor:
522 """C[M,N] = A_fp8[M,K] @ B_fp4[N,K]^T.
523 A has per-128 act scale; B has per-32 E8M0 weight scale.
524 B is stored as [N, K//2] in float4_e2m1fn_x2 (2 FP4 values per byte, packed along K)."""
525 assert a.is_contiguous() and b.is_contiguous(), "Input tensors must be contiguous"
526 assert a_s.is_contiguous() and b_s.is_contiguous(), (
527 "Scaling factor tensors must be contiguous"
528 )
529 tl_dtype = FE8M0 if scale_dtype == torch.float8_e8m0fnu else FP32
530 K = a.size(-1)
531 M = a.numel() // K
532 N = b.size(0)
533 c = a.new_empty(*a.size()[:-1], N, dtype=torch.get_default_dtype())
534 kernel = fp4_gemm_kernel(N, K, scale_dtype=tl_dtype)
535 kernel(a.view(M, K), b, c.view(M, N), a_s.view(M, -1), b_s)
536 return c
537