class FusedWoAKernel(VllmCuTeDSLJitKernel["FusedWoAKernel.CompileKey"]):
@dataclass(frozen=True)
class CompileKey:
# T=1..96 uses at most 46 keys per (n_groups, heads_per_group, o_lora_rank).
tokens_per_tile: int
token_tiles: int
full_tiles: bool
n_groups: int
heads_per_group: int
o_lora_rank: int
@staticmethod
def kernel(compile_key: CompileKey) -> Any:
tokens = compile_key.tokens_per_tile
token_tiles = compile_key.token_tiles
full_tiles = compile_key.full_tiles
acc_cols = 1 << (tokens - 1).bit_length()
tile_m = max(8, acc_cols)
tile_n = 128
tile_k = 128
# tcgen05 MXFP8 MMA consumes K = 32, one MX block, per instruction.
mma_k = 32
# Each group's heads split K across one cluster, one head per CTA.
heads_per_group = compile_key.heads_per_group
tiles_per_group = compile_key.o_lora_rank // tile_n
n_tiles = compile_key.n_groups * tiles_per_group
# Each CTA keeps one whole head of K in SMEM, one TMA box per stage.
num_stages = _HEAD_DIM // tile_k
w_stage_bytes = tile_n * tile_k
nope_dim = _HEAD_DIM - _ROPE_DIM
threads_per_token = _HEAD_DIM // 4
sf_col = max(16, tile_m)
tmem_cols = sf_col * 2
threads = min(acc_cols * threads_per_token, 512)
tokens_per_iter = threads // threads_per_token
@cute.kernel
def device_kernel(
x: cute.Tensor,
positions: cute.Tensor,
rope: cute.Tensor,
w: cpasync.TmaInfo,
ws: cpasync.TmaInfo,
q: cute.Tensor,
qs: cute.Tensor,
num_tokens: Int32,
):
total_tokens = num_tokens
if cutlass.const_expr(full_tiles):
total_tokens = Int32(tokens * token_tiles)
tid, _, _ = cute.arch.thread_idx()
bid, token_tile, _ = cute.arch.block_idx()
token_start = Int32(0)
if cutlass.const_expr(token_tiles > 1):
token_start = token_tile * tokens
warp = cute.arch.make_warp_uniform(tid // 32)
lane = tid % 32
split = bid % heads_per_group
tile = bid // heads_per_group
group = tile // tiles_per_group
smem = utils.SmemAllocator()
sW = smem.allocate_tensor(
Float8E4M3FN,
w.smem_layout.outer,
byte_alignment=128,
swizzle=w.smem_layout.inner,
)
sX = smem.allocate_tensor(
Float8E4M3FN,
cute.make_layout(
(tile_m, tile_k, num_stages), stride=(tile_k, 1, tile_m * tile_k)
),
byte_alignment=1024,
swizzle=cute.make_swizzle(3, 4, 3),
)
# tcgen05.cp 32x128b.warpx4 order: row r at byte (r % 32) * 16 +
# (r // 32) * 4, one 4-byte cell of MX-block scales per row.
sW_SF = smem.allocate_tensor(
Int32,
cute.make_layout(((32, 4), num_stages), stride=((4, 1), tile_n)),
128,
)
sW_SF_raw = smem.allocate_tensor(
Int32, cute.make_layout((tile_n, num_stages)), 128
)
sX_SF = smem.allocate_tensor(
Uint8,
cute.make_layout(
((32, 4), tile_k // mma_k, num_stages),
stride=((16, 4), 1, tile_n * 4),
),
128,
)
copy_fp8x4 = cute.make_copy_atom(
cute.nvgpu.CopyUniversalOp(), Float8E4M3FN, num_bits_per_copy=32
)
partial = smem.allocate_tensor(
Float32,
cute.make_layout(
(tile_n, cute.ceil_div(tokens, heads_per_group), heads_per_group)
),
128,
)
mbar_loaded = smem.allocate_array(Int64, num_stages)
mbar_ws_loaded = smem.allocate_array(Int64, 1)
mbar_done = smem.allocate_array(Int64, 1)
mbar_reduced = smem.allocate_array(Int64, 1)
taddr = smem.allocate(Int32, 4)
if tid == 0:
for stage in cutlass.range_constexpr(num_stages):
# Wait for weight TMA and this stage's quantization threads.
cute.arch.mbarrier_init(mbar_loaded + stage, 1 + threads // 4)
cute.arch.mbarrier_init(mbar_ws_loaded, 1)
cute.arch.mbarrier_init(mbar_done, 1)
cute.arch.mbarrier_init(mbar_reduced, 1)
cute.arch.mbarrier_init_fence()
# Local users of the barriers (warp 0's TMA below) need the init too.
cute.arch.sync_threads()
# Paired with cluster_wait() below: peers must see mbar_reduced
# initialized before their st.async; the split hides the barrier.
cute.arch.cluster_arrive_relaxed()
if warp == 0:
# Raw DeepGEMM MN-major scales; the MMA warp transposes them to
# the UTCCP order, which TMA cannot scatter at 4-byte granularity.
ws_tiles = cute.zipped_divide(ws.tma_tensor, (tile_n, num_stages))
with cute.arch.elect_one():
mbarrier.arrive_expect_tx(mbar_ws_loaded, tile_n * num_stages * 4)
simple_tma_copy(
ws.atom,
ws_tiles[None, (tile % tiles_per_group, split, group)],
sW_SF_raw,
mbar_ws_loaded,
)
tiles = cute.zipped_divide(w.tma_tensor, (tile_n, tile_k))
for stage in cutlass.range_constexpr(num_stages):
with cute.arch.elect_one():
mbarrier.arrive_expect_tx(mbar_loaded + stage, w_stage_bytes)
simple_tma_copy(
w.atom,
tiles[None, (tile, split * num_stages + stage)],
sW[None, None, stage],
mbar_loaded + stage,
)
# x and positions come from the PDL predecessor; everything above
# only touches weights. Padded sX rows only feed unread MMA columns.
cute.arch.griddepcontrol_wait()
# TMEM is allocated at runtime, and a PDL predecessor on this SM may
# hold it until it exits; allocating earlier can deadlock.
if warp == 0:
cute.arch.alloc_tmem(tmem_cols, taddr)
cute.arch.relinquish_tmem_alloc_permit()
copy_x = cute.make_copy_atom(
cute.nvgpu.CopyUniversalOp(), BFloat16, num_bits_per_copy=64
)
n_iters = cute.ceil_div(tokens, tokens_per_iter)
qid = (tid + 32) % threads_per_token
k = qid * 4
# Tail rows repeat the last input; their output stores are masked.
# Issue every token's x load before any math so their latencies overlap.
x_bf16 = cute.make_rmem_tensor((4, n_iters), BFloat16)
for iteration in cutlass.range_constexpr(n_iters):
token = iteration * tokens_per_iter + tid // threads_per_token
if cutlass.const_expr(tokens % tokens_per_iter == 0) or token < tokens:
x_src = cute.local_tile(
x[
cute.min(token_start + token, total_tokens - 1),
group * heads_per_group + split,
None,
],
(4,),
(qid,),
)
cute.copy(copy_x, x_src, x_bf16[None, iteration])
# Overlap the RoPE loads across tokens before consuming any of them.
cos_reg = cute.make_rmem_tensor((2, n_iters), Float32)
sin_reg = cute.make_rmem_tensor((2, n_iters), Float32)
if cutlass.const_expr(n_iters > 1):
copy_rope = cute.make_copy_atom(
cute.nvgpu.CopyUniversalOp(), Float32, num_bits_per_copy=32
)
if k >= nope_dim:
for iteration in cutlass.range_constexpr(n_iters):
token = iteration * tokens_per_iter + tid // threads_per_token
if (
cutlass.const_expr(tokens % tokens_per_iter == 0)
or token < tokens
):
src_token = cute.min(token_start + token, total_tokens - 1)
pos = positions[src_token]
freq = (k - nope_dim) // 2
rope_row = rope[pos, None]
cos_src = cute.local_tile(rope_row, (2,), (freq // 2,))
cute.copy(copy_rope, cos_src, cos_reg[None, iteration])
sin_tile = (freq + _ROPE_DIM // 2) // 2
sin_src = cute.local_tile(rope_row, (2,), (sin_tile,))
cute.copy(copy_rope, sin_src, sin_reg[None, iteration])
for iteration in cutlass.range_constexpr(n_iters):
token = iteration * tokens_per_iter + tid // threads_per_token
if cutlass.const_expr(tokens % tokens_per_iter == 0) or token < tokens:
x_f32 = cute.make_rmem_tensor((4,), Float32)
x_f32.store(x_bf16[None, iteration].load().to(Float32))
if k >= nope_dim:
# Inverse RoPE on the interleaved pairs (k, k+1) and
# (k+2, k+3); each rope row is cos || sin.
for pair in cutlass.range_constexpr(2):
if cutlass.const_expr(n_iters > 1):
cos = cos_reg[pair, iteration]
sin = sin_reg[pair, iteration]
else:
pos = positions[
cute.min(token_start + token, total_tokens - 1)
]
freq = (k - nope_dim) // 2 + pair
cos = Float32(rope[pos, freq])
sin = Float32(rope[pos, freq + _ROPE_DIM // 2])
even, odd = x_f32[2 * pair], x_f32[2 * pair + 1]
x_f32[2 * pair] = _rope_fma(even, cos, odd, sin)
x_f32[2 * pair + 1] = _rope_fma(
odd, cos, even, sin, subtract=True
)
amax = cute.arch.fmax(
cute.arch.fmax(cute.abs(x_f32[0]), cute.abs(x_f32[1])),
cute.arch.fmax(cute.abs(x_f32[2]), cute.abs(x_f32[3])),
)
amax = cute.arch.warp_reduction_max(amax, threads_in_group=8)
exponent, inv = _scale(cute.arch.fmax(amax, Float32(1e-10)))
x_fp8 = cute.make_rmem_tensor((4,), Float8E4M3FN)
x_fp8.store((x_f32.load() * inv).to(Float8E4M3FN))
cute.copy(
copy_fp8x4,
x_fp8,
cute.local_tile(
sX[token, None, k // tile_k], (4,), (qid % 32,)
),
)
if lane % 8 == 0:
sX_SF[token, k % tile_k // mma_k, k // tile_k] = Uint8(exponent)
cute.arch.fence_proxy("async.shared", space="cta")
mbarrier.arrive(mbar_loaded + qid // 32, order="release")
if warp == 0:
base = cute.make_tensor(taddr, cute.make_layout(1))[0]
# sW / sX are 128B-swizzled K-major; one stage's K fills the
# swizzle atom, so the leading byte offset is unused.
sdesc = _tcgen05.make_sdesc_128B_swizzle(LBO=0)
# Unswizzled 32 x 16 B scale tiles: SBO = one 8-row core matrix
# (16 B units at bit 32); bit 46 is the SM100 descriptor version.
sf_sbo = 8 * 16
sfdesc = Uint64((sf_sbo >> 4 << 32) | (1 << 46))
# M = tile_n weight rows (A), N = tile_m token rows (B).
idesc = _tcgen05.make_mxfp8_idesc(tile_n, tile_m)
# sW_SF's layout is the UTCCP order, so a plain copy transposes;
# fence the generic-proxy writes before tcgen05.cp reads them.
cute.arch.mbarrier_wait(mbar_ws_loaded, 0)
for stage in cutlass.range_constexpr(num_stages):
for i in cutlass.range_constexpr(4):
sW_SF[i * 32 + lane, stage] = sW_SF_raw[i * 32 + lane, stage]
cute.arch.fence_proxy("async.shared", space="cta")
cute.arch.sync_warp()
for stage in cutlass.range_constexpr(num_stages):
cute.arch.mbarrier_wait(mbar_loaded + stage, 0)
_tcgen05.fence_after_thread_sync()
if cutlass.const_expr(stage == num_stages - 1):
with cute.arch.elect_one():
cute.arch.griddepcontrol_launch_dependents()
adesc = sdesc | (sW[None, None, stage].iterator.toint() >> 4)
bdesc = sdesc | (sX[None, None, stage].iterator.toint() >> 4)
_tcgen05.cp(
base + sf_col,
sfdesc | (sW_SF[None, stage].iterator.toint() >> 4),
"32x128b",
"warpx4",
)
_tcgen05.cp(
base + sf_col + 4,
sfdesc | (sX_SF[None, None, stage].iterator.toint() >> 4),
"32x128b",
"warpx4",
)
# Advance mma_k FP8 bytes (16 B units) and pick MX block kk of
# each 4-byte scale (B sf_id at bit 4, A sf_id at bit 29).
for kk in cutlass.range_constexpr(tile_k // mma_k):
_tcgen05.mma_mxfp8(
base,
adesc + kk * mma_k // 16,
bdesc + kk * mma_k // 16,
idesc + ((kk << 4) | (kk << 29)),
base + sf_col,
base + sf_col + 4,
stage > 0 or kk > 0,
)
_tcgen05.commit(mbar_done)
# Only the MMA warp waits; barrier 2 releases the TMEM readers.
cute.arch.mbarrier_wait(mbar_done, 0)
if tid < tile_n:
cute.arch.barrier(barrier_id=2, number_of_threads=tile_n)
# Tokens t with t % heads_per_group == split arrive from every peer
# (this CTA included), tile_n Float32 rows each.
if tid == 0:
owned = cute.ceil_div(tokens - split, heads_per_group)
mbarrier.arrive_expect_tx(
mbar_reduced, owned * tile_n * heads_per_group * 4
)
# Pairs with cluster_arrive_relaxed() above: every peer's
# mbar_reduced is initialized before the st.async below.
cute.arch.cluster_wait()
# quack's process-wide const_expr-if rewrite can leave these unbound
# after the dynamic path above; bind them so the regions join.
amax, exponent, inv = Float32(0), Uint32(0), Float32(0)
if tid < tile_n:
base = cute.make_tensor(taddr, cute.make_layout(1))[0]
_tcgen05.fence_after_thread_sync()
acc = cute.make_rmem_tensor(acc_cols, Float32)
if cutlass.const_expr(tokens == 1):
acc[0] = _tcgen05.ld(warp * 32, base, "32x32b", 1)
else:
acc.store(_tcgen05.ld(warp * 32, base, "32x32b", acc_cols))
_tcgen05.wait_ld()
for token in cutlass.range_constexpr(tokens):
ptr = cute.domain_offset(
(tid, token // heads_per_group, split), partial
).iterator
cute.arch.store_async_dsmem(
ptr,
recast_val(acc[token], Int32),
mbar_reduced,
token % heads_per_group,
)
if split < tokens:
# One warp waits for the peers' bytes; named barrier 1 (0 is
# sync_threads) releases the other tile_n epilogue threads.
if warp == 0:
cute.arch.mbarrier_wait(mbar_reduced, 0)
cute.arch.barrier(barrier_id=1, number_of_threads=tile_n)
for i in cutlass.range_constexpr(
cute.ceil_div(tokens, heads_per_group)
):
token = split + i * heads_per_group
if cutlass.const_expr(
full_tiles
and (
tokens <= heads_per_group
or tokens % heads_per_group == 0
)
) or (token < tokens and token_start + token < total_tokens):
value = Float32(0)
for peer in cutlass.range_constexpr(heads_per_group):
value += partial[tid, i, peer]
value = Float32(BFloat16(value))
amax = cute.arch.warp_redux_sync(value, "max", abs=True)
exponent, inv = _scale(amax)
q[token_start + token, tile * tile_n + tid] = Float8E4M3FN(
value * inv
)
if lane == 0:
row = token_start + token
offset = row % 32 * 16 + row // 32 * 4 + warp
qs[tile * 512 + offset] = Uint8(exponent)
# Every scale tile owns its padding, with no second memset kernel.
for i in cutlass.range_constexpr(4):
offset = i * tile_n + tid
row = offset // 16 + (offset % 16 // 4) * 32
if token_tile == 0 and split == 0 and row >= total_tokens:
qs[tile * 512 + offset] = Uint8(0)
# Each receiver drained all writes into its own inbox. Only local
# TMEM readers must finish before deallocation; no peer reads SMEM.
if tid < tile_n:
cute.arch.barrier(barrier_id=2, number_of_threads=tile_n)
if warp == 0:
cute.arch.dealloc_tmem(
cute.make_ptr(
Float32,
cute.make_tensor(taddr, cute.make_layout(1))[0],
cute.AddressSpace.tmem,
),
tmem_cols,
)
@cute.jit
def host_entrypoint(
x: cute.Tensor,
positions: cute.Tensor,
rope: cute.Tensor,
w: cute.Tensor,
ws: cute.Tensor,
q: cute.Tensor,
qs: cute.Tensor,
num_tokens: Int32,
stream: CUstream,
):
layout = cute.make_composed_layout(
cute.make_swizzle(3, 4, 3),
0,
cute.make_layout(
(tile_n, tile_k, num_stages), stride=(tile_k, 1, w_stage_bytes)
),
)
tma = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileG2SOp(cta_group=tcgen05.CtaGroup.ONE),
w,
layout,
(tile_n, tile_k),
)
ws_tma = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileG2SOp(cta_group=tcgen05.CtaGroup.ONE),
cute.make_tensor(ws.iterator, cute.select(ws.layout, mode=[1, 2, 0])),
cute.make_layout((tile_n, num_stages)),
(tile_n, num_stages),
)
device_kernel(x, positions, rope, tma, ws_tma, q, qs, num_tokens).launch(
grid=(n_tiles * heads_per_group, token_tiles, 1),
block=(threads, 1, 1),
cluster=(heads_per_group, 1, 1),
stream=stream,
use_pdl=True,
)
return host_entrypoint
def dispatch( # type: ignore[override]
self, *, tokens: int, n_groups: int, heads_per_group: int, o_lora_rank: int
) -> CompileKey:
token_tiles = (tokens + 31) // 32
# Host integer arithmetic: balance tiles in multiples of four tokens.
tile_tokens = (
tokens
if token_tiles == 1
else (tokens + 4 * token_tiles - 1) // (4 * token_tiles) * 4
)
return self.CompileKey(
tokens_per_tile=tile_tokens,
token_tiles=token_tiles,
full_tiles=tokens % tile_tokens == 0,
n_groups=n_groups,
heads_per_group=heads_per_group,
o_lora_rank=o_lora_rank,
)
def get_warmup_keys(
self, *, max_tokens: int, n_groups: int, heads_per_group: int, o_lora_rank: int
) -> list[CompileKey]:
# Trace all accepted token counts and deduplicate shared tile shapes.
return self._trace_dispatch(self.dispatch)(
tokens=WarmupIntRange(1, max_tokens + 1),
n_groups=n_groups,
heads_per_group=heads_per_group,
o_lora_rank=o_lora_rank,
)
def warmup_inputs(self, compile_key: CompileKey) -> tuple[Any, ...]:
tokens = (
compile_key.tokens_per_tile * compile_key.token_tiles
if compile_key.full_tiles
else cute.sym_int()
)
groups = compile_key.n_groups
rank = compile_key.o_lora_rank
n = groups * rank
k = compile_key.heads_per_group * _HEAD_DIM
return (
make_fake_tensor(
BFloat16,
(tokens, groups * compile_key.heads_per_group, _HEAD_DIM),
(cute.sym_int(divisibility=8), _HEAD_DIM, 1),
assumed_align=16,
),
make_fake_tensor(Int64, (tokens,), (1,)),
make_fake_tensor(Float32, (cute.sym_int(), _ROPE_DIM), (_ROPE_DIM, 1)),
make_fake_tensor(Float8E4M3FN, (n, k), (k, 1), assumed_align=16),
make_fake_tensor(
Int32,
(groups, rank, k // 128),
(rank * k // 128, 1, rank),
assumed_align=16,
),
make_fake_tensor(Float8E4M3FN, (tokens, n), (n, 1)),
make_fake_tensor(Uint8, (128 * n // 32,), (1,)),
Int32(0),
)
@kernel_launcher
def __call__(
self,
*,
x: torch.Tensor,
positions: torch.Tensor,
rope: torch.Tensor,
weight: torch.Tensor,
weight_scale: torch.Tensor,
) -> CuTeDSLLaunchSpec["FusedWoAKernel.CompileKey"]:
"""Return FP8 WO-A output and FlashInfer F8_128x4 scale bytes."""
tokens = x.shape[0]
groups, rank = weight_scale.shape[:2]
n = groups * rank
q = torch.empty((tokens, n), device=x.device, dtype=torch.float8_e4m3fn)
scales = torch.empty(128 * n // 32, device=x.device, dtype=torch.uint8)
launch_args = (
x,
positions,
rope,
weight.view(n, -1),
weight_scale,
q,
scales,
tokens,
)
compile_key = self.dispatch(
tokens=tokens,
n_groups=groups,
heads_per_group=x.shape[1] // groups,
o_lora_rank=rank,
)
return compile_key, launch_args, (q, scales)