From 7218a144ef9f527bb474c2c8cddb4ebb77b2a962 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Tue, 11 Aug 2026 14:03:20 +0800 Subject: [PATCH 01/30] feat(distributed): add SM90 FP8 mega MoE --- .../mega_moe/example_sm90_fp8_mega_moe.py | 739 ++++++++++++++++++ .../test_example_sm90_fp8_mega_moe.py | 30 + 2 files changed, 769 insertions(+) create mode 100644 examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py create mode 100644 examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py new file mode 100644 index 0000000000..af0f429446 --- /dev/null +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -0,0 +1,739 @@ +"""Multi-GPU FP8 Mega MoE for SM90 using TileScale distributed primitives. + +This correctness-first implementation uses symmetric VMM buffers for expert +dispatch and combine. The two FP8 GEMMs use per-token/per-128 activation scales +and per-(128, 128) weight scales. Host barriers separate the communication +phases for now; later variants can overlap those phases with persistent kernels. +""" + +from __future__ import annotations + +import argparse +import math +import os +from typing import Tuple + +import torch +import torch.distributed as dist +import torch.multiprocessing + +import tilelang +import tilelang.language as T +from tilelang.distributed.allocator import get_allocator +from tilelang.distributed.bench import do_bench +from tilelang.distributed.host import init_dist + +os.environ.setdefault("NCCL_DEBUG", "ERROR") + + +MODEL_CONFIGS = { + "smoke": {"hidden": 512, "intermediate_hidden": 512, "num_experts": 8, "num_topk": 2}, + "flash": {"hidden": 4096, "intermediate_hidden": 2048, "num_experts": 256, "num_topk": 6}, + "pro": {"hidden": 7168, "intermediate_hidden": 3072, "num_experts": 384, "num_topk": 6}, +} + +FP8_MAX = 448.0 +SCALE_GRANULARITY = 128 + + +def ceil_div(x: int, y: int) -> int: + return (x + y - 1) // y + + +def align_up(x: int, alignment: int) -> int: + return ceil_div(x, alignment) * alignment + + +def per_token_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: + m, k = x.shape + x_view = x.float().view(m, k // SCALE_GRANULARITY, SCALE_GRANULARITY) + amax = x_view.abs().amax(dim=-1).clamp(1e-4) + scale = amax / FP8_MAX + x_fp8 = (x_view / scale.unsqueeze(-1)).to(torch.float8_e4m3fn) + return x_fp8.view(m, k).contiguous(), scale.contiguous() + + +def block_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: + groups, n, k = x.shape + x_view = x.float().view( + groups, + n // SCALE_GRANULARITY, + SCALE_GRANULARITY, + k // SCALE_GRANULARITY, + SCALE_GRANULARITY, + ) + amax = x_view.abs().amax(dim=(-1, -3)).clamp(1e-4) + scale = amax / FP8_MAX + x_fp8 = (x_view / scale.unsqueeze(-1).unsqueeze(-3)).to(torch.float8_e4m3fn) + return x_fp8.view(groups, n, k).contiguous(), scale.contiguous() + + +def dequantize_per_token(x: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: + m, k = x.shape + return (x.float().view(m, k // SCALE_GRANULARITY, SCALE_GRANULARITY) * scale.unsqueeze(-1)).view(m, k) + + +def dequantize_block(x: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: + groups, n, k = x.shape + x_view = x.float().view( + groups, + n // SCALE_GRANULARITY, + SCALE_GRANULARITY, + k // SCALE_GRANULARITY, + SCALE_GRANULARITY, + ) + return (x_view * scale.unsqueeze(-1).unsqueeze(-3)).view(groups, n, k) + + +def assign_local_routes_kernel( + num_tokens: int, + num_experts: int, + num_topk: int, + num_ranks: int, + threads: int = 128, +): + @T.prim_func + def main( + topk_idx: T.Tensor((num_tokens, num_topk), T.int32), + route_counts: T.Tensor((num_ranks, num_experts), T.int32), + route_slots: T.Tensor((num_tokens, num_topk), T.int32), + ): + with T.Kernel(T.ceildiv(num_tokens * num_topk, threads), threads=threads) as bx: + src_rank = T.alloc_local((1,), T.int32) + src_rank[0] = T.get_rank() + route_idx = bx * threads + T.get_thread_binding() + if src_rank[0] < num_ranks and route_idx < num_tokens * num_topk: + token_idx = route_idx // num_topk + topk_slot = route_idx % num_topk + expert_idx = topk_idx[token_idx, topk_slot] + if expert_idx >= 0 and expert_idx < num_experts: + route_slots[token_idx, topk_slot] = T.atomic_add( + route_counts[src_rank[0], expert_idx], + 1, + memory_order="relaxed", + return_prev=True, + ) + else: + route_slots[token_idx, topk_slot] = -1 + + return main + + +def publish_route_counts_kernel(num_experts: int, num_ranks: int, threads: int = 128): + @T.prim_func + def main(route_counts: T.Tensor((num_ranks, num_experts), T.int32)): + with T.Kernel(T.ceildiv(num_experts * num_ranks, threads), threads=threads) as bx: + src_rank = T.alloc_local((1,), T.int32) + src_rank[0] = T.get_rank() + idx = bx * threads + T.get_thread_binding() + if idx < num_experts * num_ranks: + dst_rank = idx // num_experts + expert_idx = idx % num_experts + if dst_rank != src_rank[0]: + T.st( + route_counts[src_rank[0], expert_idx], + route_counts[src_rank[0], expert_idx], + dst_pe=dst_rank, + ) + + return main + + +def finalize_routes_kernel( + num_tokens: int, + hidden: int, + num_experts: int, + num_topk: int, + num_ranks: int, + capacity: int, + threads: int = 128, +): + num_experts_per_rank = num_experts // num_ranks + num_scale_groups = hidden // SCALE_GRANULARITY + num_routes = num_tokens * num_topk + num_work_items = max(num_routes, num_experts_per_rank) + + @T.prim_func + def main( + x_sf: T.Tensor((num_tokens, num_scale_groups), T.float32), + topk_idx: T.Tensor((num_tokens, num_topk), T.int32), + topk_weights: T.Tensor((num_tokens, num_topk), T.float32), + route_counts: T.Tensor((num_ranks, num_experts), T.int32), + recv_counts: T.Tensor((num_experts_per_rank,), T.int32), + recv_x_sf: T.Tensor((num_experts_per_rank, capacity, num_scale_groups), T.float32), + recv_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), + src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), + src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), + src_topk: T.Tensor((num_experts_per_rank, capacity), T.int32), + route_slots: T.Tensor((num_tokens, num_topk), T.int32), + ): + with T.Kernel(T.ceildiv(num_work_items, threads), threads=threads) as bx: + src_rank = T.alloc_local((1,), T.int32) + src_rank[0] = T.get_rank() + work_idx = bx * threads + T.get_thread_binding() + + if work_idx < num_experts_per_rank: + count = T.alloc_var(T.int32, init=0) + expert_idx = src_rank[0] * num_experts_per_rank + work_idx + for peer_rank in T.serial(num_ranks): + count += route_counts[peer_rank, expert_idx] + recv_counts[work_idx] = count + + if work_idx < num_routes: + token_idx = work_idx // num_topk + topk_slot = work_idx % num_topk + expert_idx = topk_idx[token_idx, topk_slot] + slot = T.alloc_var(T.int32, init=route_slots[token_idx, topk_slot]) + if token_idx < num_tokens and expert_idx >= 0 and expert_idx < num_experts and slot >= 0: + for peer_rank in T.serial(num_ranks): + if peer_rank < src_rank[0]: + slot += route_counts[peer_rank, expert_idx] + route_slots[token_idx, topk_slot] = slot + if slot < capacity: + dst_rank = expert_idx // num_experts_per_rank + local_expert = expert_idx % num_experts_per_rank + T.st(recv_weights[local_expert, slot], topk_weights[token_idx, topk_slot], dst_pe=dst_rank) + T.st(src_ranks[local_expert, slot], src_rank[0], dst_pe=dst_rank) + T.st(src_tokens[local_expert, slot], token_idx, dst_pe=dst_rank) + T.st(src_topk[local_expert, slot], topk_slot, dst_pe=dst_rank) + for scale_idx in T.serial(num_scale_groups): + T.st( + recv_x_sf[local_expert, slot, scale_idx], + x_sf[token_idx, scale_idx], + dst_pe=dst_rank, + ) + + return main + + +def dispatch_tokens_kernel( + num_tokens: int, + hidden: int, + num_experts: int, + num_topk: int, + num_ranks: int, + capacity: int, + block_h: int = 256, + threads: int = 128, +): + num_experts_per_rank = num_experts // num_ranks + + @T.prim_func + def main( + x: T.Tensor((num_tokens, hidden), T.float8_e4m3fn), + topk_idx: T.Tensor((num_tokens, num_topk), T.int32), + route_slots: T.Tensor((num_tokens, num_topk), T.int32), + recv_x: T.Tensor((num_experts_per_rank, capacity, hidden), T.float8_e4m3fn), + ): + with T.Kernel(T.ceildiv(hidden, block_h), num_tokens * num_topk, threads=threads) as (bx, by): + token_idx = by // num_topk + topk_slot = by % num_topk + expert_idx = topk_idx[token_idx, topk_slot] + slot = route_slots[token_idx, topk_slot] + if expert_idx >= 0 and slot >= 0 and slot < capacity: + dst_rank = expert_idx // num_experts_per_rank + local_expert = expert_idx % num_experts_per_rank + T.copy( + x[token_idx, bx * block_h : (bx + 1) * block_h], + recv_x[local_expert, slot, bx * block_h : (bx + 1) * block_h], + dst_pe=dst_rank, + disable_tma=True, + ) + T.fence_sys() + + return main + + +def fp8_grouped_gemm_kernel( + num_experts_per_rank: int, + capacity: int, + n: int, + k: int, + block_m: int = 64, + block_n: int = 128, + block_k: int = 128, + threads: int = 128, + pipeline_stages: int = 4, +): + @T.prim_func + def main( + a: T.Tensor((num_experts_per_rank, capacity, k), T.float8_e4m3fn), + b: T.Tensor((num_experts_per_rank, n, k), T.float8_e4m3fn), + a_sf: T.Tensor((num_experts_per_rank, capacity, k // SCALE_GRANULARITY), T.float32), + b_sf: T.Tensor( + (num_experts_per_rank, n // SCALE_GRANULARITY, k // SCALE_GRANULARITY), + T.float32, + ), + recv_counts: T.Tensor((num_experts_per_rank,), T.int32), + out: T.Tensor((num_experts_per_rank, capacity, n), T.bfloat16), + ): + with T.Kernel(T.ceildiv(n, block_n), T.ceildiv(capacity, block_m), num_experts_per_rank, threads=threads) as ( + bx, + by, + bz, + ): + a_shared = T.alloc_shared((block_m, block_k), T.float8_e4m3fn) + b_shared = T.alloc_shared((block_n, block_k), T.float8_e4m3fn) + out_shared = T.alloc_shared((block_m, block_n), T.bfloat16) + partial = T.alloc_fragment((block_m, block_n), T.float32) + accum = T.alloc_fragment((block_m, block_n), T.float32) + + if by * block_m < recv_counts[bz]: + T.clear(partial) + T.clear(accum) + for ko in T.Pipelined(k // block_k, num_stages=pipeline_stages): + T.copy(a[bz, by * block_m, ko * block_k], a_shared) + T.copy(b[bz, bx * block_n, ko * block_k], b_shared) + T.gemm(a_shared, b_shared, partial, transpose_B=True) + b_scale = b_sf[bz, bx, ko] + for i, j in T.Parallel(block_m, block_n): + accum[i, j] += partial[i, j] * (a_sf[bz, by * block_m + i, ko] * b_scale) + T.clear(partial) + T.copy(accum, out_shared) + T.copy(out_shared, out[bz, by * block_m, bx * block_n]) + + return main + + +def swiglu_quant_kernel( + num_experts_per_rank: int, + capacity: int, + intermediate_hidden: int, + block_m: int = 8, + block_n: int = 128, + threads: int = 128, + activation_clamp: float = 10.0, +): + @T.prim_func + def main( + gate_up: T.Tensor((num_experts_per_rank, capacity, 2 * intermediate_hidden), T.bfloat16), + route_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), + recv_counts: T.Tensor((num_experts_per_rank,), T.int32), + out: T.Tensor((num_experts_per_rank, capacity, intermediate_hidden), T.float8_e4m3fn), + out_sf: T.Tensor( + (num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), + T.float32, + ), + ): + with T.Kernel( + intermediate_hidden // block_n, + T.ceildiv(capacity, block_m), + num_experts_per_rank, + threads=threads, + ) as (bx, by, bz): + gate = T.alloc_fragment((block_m, block_n), T.float32) + up = T.alloc_fragment((block_m, block_n), T.float32) + activated = T.alloc_fragment((block_m, block_n), T.float32) + amax = T.alloc_fragment((block_m,), T.float32) + scale = T.alloc_fragment((block_m,), T.float32) + quant = T.alloc_fragment((block_m, block_n), T.float32) + quant_fp8 = T.alloc_fragment((block_m, block_n), T.float8_e4m3fn) + + if by * block_m < recv_counts[bz]: + T.copy(gate_up[bz, by * block_m, bx * block_n], gate) + T.copy(gate_up[bz, by * block_m, intermediate_hidden + bx * block_n], up) + for i, j in T.Parallel(block_m, block_n): + gate[i, j] = T.min(gate[i, j], activation_clamp) + up[i, j] = T.max(T.min(up[i, j], activation_clamp), -activation_clamp) + activated[i, j] = ( + gate[i, j] + * T.sigmoid(gate[i, j]) + * up[i, j] + * route_weights[bz, by * block_m + i] + ) + T.reduce_absmax(activated, amax, dim=1) + for i in T.Parallel(block_m): + scale[i] = T.max(amax[i], 1e-4) / FP8_MAX + out_sf[bz, by * block_m + i, bx] = scale[i] + for i, j in T.Parallel(block_m, block_n): + quant[i, j] = T.clamp(activated[i, j] / scale[i], -FP8_MAX, FP8_MAX) + T.copy(quant, quant_fp8) + T.copy(quant_fp8, out[bz, by * block_m, bx * block_n]) + + return main + + +def scatter_outputs_kernel( + num_experts_per_rank: int, + capacity: int, + num_tokens: int, + num_topk: int, + hidden: int, + block_h: int = 256, + threads: int = 128, +): + @T.prim_func + def main( + local_out: T.Tensor((num_experts_per_rank, capacity, hidden), T.bfloat16), + recv_counts: T.Tensor((num_experts_per_rank,), T.int32), + src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), + src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), + src_topk: T.Tensor((num_experts_per_rank, capacity), T.int32), + combine: T.Tensor((num_tokens, num_topk, hidden), T.bfloat16), + ): + with T.Kernel(T.ceildiv(hidden, block_h), capacity, num_experts_per_rank, threads=threads) as (bx, by, bz): + if by < recv_counts[bz]: + dst_rank = src_ranks[bz, by] + token_idx = src_tokens[bz, by] + topk_slot = src_topk[bz, by] + if dst_rank >= 0 and token_idx < num_tokens and topk_slot < num_topk: + T.copy( + local_out[bz, by, bx * block_h : (bx + 1) * block_h], + combine[token_idx, topk_slot, bx * block_h : (bx + 1) * block_h], + dst_pe=dst_rank, + disable_tma=True, + ) + T.fence_sys() + + return main + + +def reduce_topk_kernel( + num_tokens: int, + num_topk: int, + hidden: int, + block_m: int = 8, + block_h: int = 128, + threads: int = 128, +): + @T.prim_func + def main( + combine: T.Tensor((num_tokens, num_topk, hidden), T.bfloat16), + out: T.Tensor((num_tokens, hidden), T.bfloat16), + ): + with T.Kernel(T.ceildiv(hidden, block_h), T.ceildiv(num_tokens, block_m), threads=threads) as (bx, by): + accum = T.alloc_fragment((block_m, block_h), T.float32) + out_shared = T.alloc_shared((block_m, block_h), T.bfloat16) + T.clear(accum) + for topk_slot in T.serial(num_topk): + for i, j in T.Parallel(block_m, block_h): + if by * block_m + i < num_tokens: + accum[i, j] += combine[by * block_m + i, topk_slot, bx * block_h + j] + T.copy(accum, out_shared) + T.copy(out_shared, out[by * block_m, bx * block_h]) + + return main + + +def _allocator_size_bytes( + num_tokens: int, + hidden: int, + intermediate_hidden: int, + num_experts_per_rank: int, + num_topk: int, + capacity: int, +) -> int: + fp8 = 1 + bf16 = 2 + fp32 = 4 + i32 = 4 + weight_bytes = num_experts_per_rank * ( + 2 * intermediate_hidden * hidden * fp8 + hidden * intermediate_hidden * fp8 + ) + weight_scale_bytes = num_experts_per_rank * ( + (2 * intermediate_hidden // 128) * (hidden // 128) + + (hidden // 128) * (intermediate_hidden // 128) + ) * fp32 + pool_bytes = num_experts_per_rank * capacity * ( + hidden * fp8 + + (hidden // 128) * fp32 + + 4 * i32 + + 2 * intermediate_hidden * bf16 + + intermediate_hidden * fp8 + + (intermediate_hidden // 128) * fp32 + + hidden * bf16 + ) + input_bytes = num_tokens * ( + hidden * fp8 + (hidden // 128) * fp32 + num_topk * (3 * i32 + fp32) + num_topk * hidden * bf16 + ) + return align_up(weight_bytes + weight_scale_bytes + pool_bytes + input_bytes + 2**27, 2**20) + + +def _gather_cat(tensor: torch.Tensor, group: dist.ProcessGroup) -> torch.Tensor: + gathered = [torch.empty_like(tensor) for _ in range(dist.get_world_size(group))] + dist.all_gather(gathered, tensor, group=group) + return torch.cat(gathered, dim=0) + + +def torch_reference( + x_fp8: torch.Tensor, + x_sf: torch.Tensor, + topk_idx: torch.Tensor, + topk_weights: torch.Tensor, + l1_fp8: torch.Tensor, + l1_sf: torch.Tensor, + l2_fp8: torch.Tensor, + l2_sf: torch.Tensor, + group: dist.ProcessGroup, + activation_clamp: float, +) -> torch.Tensor: + l1_all = _gather_cat(l1_fp8, group) + l1_sf_all = _gather_cat(l1_sf, group) + l2_all = _gather_cat(l2_fp8, group) + l2_sf_all = _gather_cat(l2_sf, group) + x = dequantize_per_token(x_fp8, x_sf) + result = torch.zeros((x.size(0), l2_all.size(1)), dtype=torch.float32, device=x.device) + + for expert_idx in range(l1_all.size(0)): + positions = (topk_idx == expert_idx).nonzero(as_tuple=False) + if positions.numel() == 0: + continue + token_indices = positions[:, 0] + topk_slots = positions[:, 1] + l1_weight = dequantize_block(l1_all[expert_idx : expert_idx + 1], l1_sf_all[expert_idx : expert_idx + 1])[0] + gate_up = x[token_indices] @ l1_weight.T + gate, up = gate_up.chunk(2, dim=-1) + gate = gate.clamp(max=activation_clamp) + up = up.clamp(min=-activation_clamp, max=activation_clamp) + activated = torch.nn.functional.silu(gate) * up + activated *= topk_weights[token_indices, topk_slots].unsqueeze(-1) + activated_fp8, activated_sf = per_token_cast_to_fp8(activated) + activated_dequant = dequantize_per_token(activated_fp8, activated_sf) + l2_weight = dequantize_block(l2_all[expert_idx : expert_idx + 1], l2_sf_all[expert_idx : expert_idx + 1])[0] + contribution = (activated_dequant @ l2_weight.T).to(torch.bfloat16).float() + result.index_add_(0, token_indices, contribution) + + return result.to(torch.bfloat16) + + +def calc_diff(x: torch.Tensor, y: torch.Tensor) -> float: + x, y = x.double(), y.double() + return (1 - 2 * (x * y).sum() / (x.square() + y.square()).sum()).item() + + +def allocator_tensor(shape, dtype, allocator): + if dtype == torch.float8_e4m3fn: + return tilelang.tensor(shape, torch.uint8, allocator=allocator).view(dtype) + return tilelang.tensor(shape, dtype, allocator=allocator) + + +def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): + model = MODEL_CONFIGS[args.model_config] + hidden = model["hidden"] + intermediate_hidden = model["intermediate_hidden"] + num_experts = model["num_experts"] + num_topk = model["num_topk"] + num_tokens = args.num_tokens + activation_clamp = args.activation_clamp + + assert num_experts % num_local_ranks == 0 + assert hidden % 256 == 0 and intermediate_hidden % 128 == 0 + num_experts_per_rank = num_experts // num_local_ranks + average_recv = ceil_div(num_tokens * num_local_ranks * num_topk, num_experts) + capacity = args.capacity or align_up(max(average_recv * 2, 64), 64) + + rank, num_ranks, group = init_dist(local_rank, num_local_ranks) + assert rank == local_rank and num_ranks == num_local_ranks + allocator = get_allocator( + size=_allocator_size_bytes( + num_tokens, + hidden, + intermediate_hidden, + num_experts_per_rank, + num_topk, + capacity, + ), + device=f"cuda:{local_rank}", + is_distributed=True, + local_rank=local_rank, + num_local_ranks=num_local_ranks, + group=group, + use_vmm=True, + ) + + kernel_specs = [ + assign_local_routes_kernel(num_tokens, num_experts, num_topk, num_ranks), + publish_route_counts_kernel(num_experts, num_ranks), + finalize_routes_kernel(num_tokens, hidden, num_experts, num_topk, num_ranks, capacity), + dispatch_tokens_kernel(num_tokens, hidden, num_experts, num_topk, num_ranks, capacity), + fp8_grouped_gemm_kernel(num_experts_per_rank, capacity, 2 * intermediate_hidden, hidden), + swiglu_quant_kernel( + num_experts_per_rank, + capacity, + intermediate_hidden, + activation_clamp=activation_clamp, + ), + fp8_grouped_gemm_kernel(num_experts_per_rank, capacity, hidden, intermediate_hidden), + scatter_outputs_kernel(num_experts_per_rank, capacity, num_tokens, num_topk, hidden), + reduce_topk_kernel(num_tokens, num_topk, hidden), + ] + kernels = [tilelang.compile(spec, compile_once=True, compile_group=group) for spec in kernel_specs] + for kernel in kernels: + kernel.initialize(allocator=allocator) + ( + assign_local_routes, + publish_route_counts, + finalize_routes, + dispatch_tokens, + l1_gemm, + swiglu_quant, + l2_gemm, + scatter_outputs, + reduce_topk, + ) = kernels + + if local_rank == 0 and args.print_source: + for kernel in kernels: + print(kernel.get_kernel_source()) + + torch.manual_seed(args.seed + local_rank) + x_bf16 = torch.randn((num_tokens, hidden), dtype=torch.bfloat16, device="cuda") + x_fp8_src, x_sf_src = per_token_cast_to_fp8(x_bf16) + scores = torch.randn((num_tokens, num_experts), dtype=torch.float32, device="cuda") + topk_weights_src, topk_idx_src = torch.topk(scores, num_topk, dim=-1, sorted=False) + topk_idx_src = topk_idx_src.to(torch.int32) + + l1_bf16 = torch.randn( + (num_experts_per_rank, 2 * intermediate_hidden, hidden), + dtype=torch.bfloat16, + device="cuda", + ) * 0.05 + l2_bf16 = torch.randn( + (num_experts_per_rank, hidden, intermediate_hidden), + dtype=torch.bfloat16, + device="cuda", + ) * 0.05 + l1_fp8_src, l1_sf_src = block_cast_to_fp8(l1_bf16) + l2_fp8_src, l2_sf_src = block_cast_to_fp8(l2_bf16) + del scores, l1_bf16, l2_bf16 + + x = allocator_tensor(x_fp8_src.shape, x_fp8_src.dtype, allocator=allocator).copy_(x_fp8_src) + x_sf = allocator_tensor(x_sf_src.shape, x_sf_src.dtype, allocator=allocator).copy_(x_sf_src) + topk_idx = allocator_tensor(topk_idx_src.shape, topk_idx_src.dtype, allocator=allocator).copy_(topk_idx_src) + topk_weights = allocator_tensor( + topk_weights_src.shape, topk_weights_src.dtype, allocator=allocator + ).copy_(topk_weights_src) + l1_fp8 = allocator_tensor(l1_fp8_src.shape, l1_fp8_src.dtype, allocator=allocator).copy_(l1_fp8_src) + l1_sf = allocator_tensor(l1_sf_src.shape, l1_sf_src.dtype, allocator=allocator).copy_(l1_sf_src) + l2_fp8 = allocator_tensor(l2_fp8_src.shape, l2_fp8_src.dtype, allocator=allocator).copy_(l2_fp8_src) + l2_sf = allocator_tensor(l2_sf_src.shape, l2_sf_src.dtype, allocator=allocator).copy_(l2_sf_src) + + route_counts = allocator_tensor((num_ranks, num_experts), torch.int32, allocator=allocator) + recv_counts = allocator_tensor((num_experts_per_rank,), torch.int32, allocator=allocator) + recv_x = allocator_tensor((num_experts_per_rank, capacity, hidden), torch.float8_e4m3fn, allocator=allocator) + recv_x_sf = allocator_tensor( + (num_experts_per_rank, capacity, hidden // SCALE_GRANULARITY), torch.float32, allocator=allocator + ) + recv_weights = allocator_tensor((num_experts_per_rank, capacity), torch.float32, allocator=allocator) + src_ranks = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) + src_tokens = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) + src_topk = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) + route_slots = allocator_tensor((num_tokens, num_topk), torch.int32, allocator=allocator) + l1_out = allocator_tensor( + (num_experts_per_rank, capacity, 2 * intermediate_hidden), torch.bfloat16, allocator=allocator + ) + l2_x = allocator_tensor( + (num_experts_per_rank, capacity, intermediate_hidden), torch.float8_e4m3fn, allocator=allocator + ) + l2_x_sf = allocator_tensor( + (num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), + torch.float32, + allocator=allocator, + ) + l2_out = allocator_tensor((num_experts_per_rank, capacity, hidden), torch.bfloat16, allocator=allocator) + combine = allocator_tensor((num_tokens, num_topk, hidden), torch.bfloat16, allocator=allocator) + out = allocator_tensor((num_tokens, hidden), torch.bfloat16, allocator=allocator) + + def reset_state(): + route_counts.zero_() + recv_counts.zero_() + recv_x.zero_() + recv_x_sf.zero_() + recv_weights.zero_() + src_ranks.fill_(-1) + combine.zero_() + torch.cuda.synchronize() + dist.barrier(group=group) + + def run_pipeline(check_capacity: bool = False): + assign_local_routes(topk_idx, route_counts, route_slots) + publish_route_counts(route_counts) + torch.cuda.synchronize() + dist.barrier(group=group) + finalize_routes( + x_sf, + topk_idx, + topk_weights, + route_counts, + recv_counts, + recv_x_sf, + recv_weights, + src_ranks, + src_tokens, + src_topk, + route_slots, + ) + dispatch_tokens(x, topk_idx, route_slots, recv_x) + torch.cuda.synchronize() + dist.barrier(group=group) + if check_capacity: + local_max = recv_counts.max() + dist.all_reduce(local_max, op=dist.ReduceOp.MAX, group=group) + assert local_max.item() <= capacity, ( + f"expert capacity {capacity} is smaller than received routes {local_max.item()}" + ) + l1_gemm(recv_x, l1_fp8, recv_x_sf, l1_sf, recv_counts, l1_out) + swiglu_quant(l1_out, recv_weights, recv_counts, l2_x, l2_x_sf) + l2_gemm(l2_x, l2_fp8, l2_x_sf, l2_sf, recv_counts, l2_out) + scatter_outputs(l2_out, recv_counts, src_ranks, src_tokens, src_topk, combine) + torch.cuda.synchronize() + dist.barrier(group=group) + reduce_topk(combine, out) + return out + + reset_state() + actual = run_pipeline(check_capacity=True) + torch.cuda.synchronize() + dist.barrier(group=group) + + if args.check: + expected = torch_reference( + x_fp8_src, + x_sf_src, + topk_idx_src, + topk_weights_src, + l1_fp8_src, + l1_sf_src, + l2_fp8_src, + l2_sf_src, + group, + activation_clamp, + ) + diff = calc_diff(actual, expected) + assert diff < args.diff_tol, f"rank {local_rank}: diff={diff} exceeds {args.diff_tol}" + print(f"rank {local_rank} check passed, diff={diff:.6f}") + + if args.rep > 0: + reset_state() + latency = do_bench( + run_pipeline, + warmup=args.warmup, + rep=args.rep, + post_fn=reset_state, + group=group, + ) + if local_rank == 0: + print( + f"tilescale sm90 fp8 mega moe: model={args.model_config} M={num_tokens} " + f"capacity={capacity} latency={latency * 1000:.1f} us" + ) + + allocator.close() + dist.destroy_process_group() + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--num-processes", type=int, default=8) + parser.add_argument("--model-config", choices=tuple(MODEL_CONFIGS), default="smoke") + parser.add_argument("--num-tokens", type=int, default=64) + parser.add_argument("--capacity", type=int, default=None) + parser.add_argument("--activation-clamp", type=float, default=10.0) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--diff-tol", type=float, default=0.01) + parser.add_argument("--warmup", type=int, default=1) + parser.add_argument("--rep", type=int, default=1) + parser.add_argument("--check", action="store_true") + parser.add_argument("--print-source", action="store_true") + args = parser.parse_args() + torch.multiprocessing.spawn(main, args=(args.num_processes, args), nprocs=args.num_processes) diff --git a/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py new file mode 100644 index 0000000000..c70f1b3e4a --- /dev/null +++ b/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +import argparse + +import tilelang.testing +from testing.python.distributed._utils import distributed_test + +import example_sm90_fp8_mega_moe + + +@distributed_test(nprocs=4, require_fabric=True) +def test_example_sm90_fp8_mega_moe(local_rank: int, num_ranks: int): + args = argparse.Namespace( + num_processes=num_ranks, + model_config="smoke", + num_tokens=32, + capacity=64, + activation_clamp=10.0, + seed=0, + diff_tol=0.01, + warmup=0, + rep=0, + check=True, + print_source=False, + ) + example_sm90_fp8_mega_moe.main(local_rank, num_ranks, args) + + +if __name__ == "__main__": + tilelang.testing.main() From dc38f2361c49277310909243aa0677969a0040be Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Tue, 11 Aug 2026 14:25:48 +0800 Subject: [PATCH 02/30] perf(distributed): use device barriers for mega MoE --- .../mega_moe/example_sm90_fp8_mega_moe.py | 49 +++++++++++++++---- 1 file changed, 39 insertions(+), 10 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index af0f429446..248cd29c08 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -1,9 +1,8 @@ """Multi-GPU FP8 Mega MoE for SM90 using TileScale distributed primitives. -This correctness-first implementation uses symmetric VMM buffers for expert -dispatch and combine. The two FP8 GEMMs use per-token/per-128 activation scales -and per-(128, 128) weight scales. Host barriers separate the communication -phases for now; later variants can overlap those phases with persistent kernels. +This implementation uses symmetric VMM buffers for expert dispatch and combine. +The two FP8 GEMMs use per-token/per-128 activation scales and per-(128, 128) +weight scales. Device-side system barriers order the communication phases. """ from __future__ import annotations @@ -119,6 +118,19 @@ def main( return main +def reset_route_counts_kernel(num_experts: int, num_ranks: int, threads: int = 128): + @T.prim_func + def main(route_counts: T.Tensor((num_ranks, num_experts), T.int32)): + with T.Kernel(T.ceildiv(num_experts, threads), threads=threads) as bx: + src_rank = T.alloc_local((1,), T.int32) + src_rank[0] = T.get_rank() + expert_idx = bx * threads + T.get_thread_binding() + if src_rank[0] < num_ranks and expert_idx < num_experts: + route_counts[src_rank[0], expert_idx] = 0 + + return main + + def publish_route_counts_kernel(num_experts: int, num_ranks: int, threads: int = 128): @T.prim_func def main(route_counts: T.Tensor((num_ranks, num_experts), T.int32)): @@ -139,6 +151,19 @@ def main(route_counts: T.Tensor((num_ranks, num_experts), T.int32)): return main +def device_barrier_kernel(num_ranks: int): + @T.prim_func + def main(barrier: T.Tensor((num_ranks,), T.int32)): + with T.Kernel(1, threads=32): + rank = T.alloc_local((1,), T.int32) + rank[0] = T.get_rank() + if rank[0] < num_ranks: + T.barrier_blocks(barrier[0]) + T.fence_sys() + + return main + + def finalize_routes_kernel( num_tokens: int, hidden: int, @@ -542,8 +567,10 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): ) kernel_specs = [ + reset_route_counts_kernel(num_experts, num_ranks), assign_local_routes_kernel(num_tokens, num_experts, num_topk, num_ranks), publish_route_counts_kernel(num_experts, num_ranks), + device_barrier_kernel(num_ranks), finalize_routes_kernel(num_tokens, hidden, num_experts, num_topk, num_ranks, capacity), dispatch_tokens_kernel(num_tokens, hidden, num_experts, num_topk, num_ranks, capacity), fp8_grouped_gemm_kernel(num_experts_per_rank, capacity, 2 * intermediate_hidden, hidden), @@ -561,8 +588,10 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): for kernel in kernels: kernel.initialize(allocator=allocator) ( + reset_route_counts, assign_local_routes, publish_route_counts, + device_barrier, finalize_routes, dispatch_tokens, l1_gemm, @@ -609,6 +638,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): l2_sf = allocator_tensor(l2_sf_src.shape, l2_sf_src.dtype, allocator=allocator).copy_(l2_sf_src) route_counts = allocator_tensor((num_ranks, num_experts), torch.int32, allocator=allocator) + barrier = allocator_tensor((num_ranks,), torch.int32, allocator=allocator) recv_counts = allocator_tensor((num_experts_per_rank,), torch.int32, allocator=allocator) recv_x = allocator_tensor((num_experts_per_rank, capacity, hidden), torch.float8_e4m3fn, allocator=allocator) recv_x_sf = allocator_tensor( @@ -636,6 +666,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): def reset_state(): route_counts.zero_() + barrier.zero_() recv_counts.zero_() recv_x.zero_() recv_x_sf.zero_() @@ -646,10 +677,10 @@ def reset_state(): dist.barrier(group=group) def run_pipeline(check_capacity: bool = False): + reset_route_counts(route_counts) assign_local_routes(topk_idx, route_counts, route_slots) publish_route_counts(route_counts) - torch.cuda.synchronize() - dist.barrier(group=group) + device_barrier(barrier) finalize_routes( x_sf, topk_idx, @@ -664,8 +695,7 @@ def run_pipeline(check_capacity: bool = False): route_slots, ) dispatch_tokens(x, topk_idx, route_slots, recv_x) - torch.cuda.synchronize() - dist.barrier(group=group) + device_barrier(barrier) if check_capacity: local_max = recv_counts.max() dist.all_reduce(local_max, op=dist.ReduceOp.MAX, group=group) @@ -676,8 +706,7 @@ def run_pipeline(check_capacity: bool = False): swiglu_quant(l1_out, recv_weights, recv_counts, l2_x, l2_x_sf) l2_gemm(l2_x, l2_fp8, l2_x_sf, l2_sf, recv_counts, l2_out) scatter_outputs(l2_out, recv_counts, src_ranks, src_tokens, src_topk, combine) - torch.cuda.synchronize() - dist.barrier(group=group) + device_barrier(barrier) reduce_topk(combine, out) return out From 3f4136d305282d0470b8a3c89a2cba35da28d7e7 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Tue, 11 Aug 2026 14:34:42 +0800 Subject: [PATCH 03/30] fix(distributed): keep mega MoE barrier offset in range --- examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 248cd29c08..51207fa713 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -626,6 +626,9 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): l2_fp8_src, l2_sf_src = block_cast_to_fp8(l2_bf16) del scores, l1_bf16, l2_bf16 + # barrier_blocks currently lowers its byte offset as int32, so keep this + # allocation before the multi-gigabyte Pro-model weights. + barrier = allocator_tensor((num_ranks,), torch.int32, allocator=allocator) x = allocator_tensor(x_fp8_src.shape, x_fp8_src.dtype, allocator=allocator).copy_(x_fp8_src) x_sf = allocator_tensor(x_sf_src.shape, x_sf_src.dtype, allocator=allocator).copy_(x_sf_src) topk_idx = allocator_tensor(topk_idx_src.shape, topk_idx_src.dtype, allocator=allocator).copy_(topk_idx_src) @@ -638,7 +641,6 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): l2_sf = allocator_tensor(l2_sf_src.shape, l2_sf_src.dtype, allocator=allocator).copy_(l2_sf_src) route_counts = allocator_tensor((num_ranks, num_experts), torch.int32, allocator=allocator) - barrier = allocator_tensor((num_ranks,), torch.int32, allocator=allocator) recv_counts = allocator_tensor((num_experts_per_rank,), torch.int32, allocator=allocator) recv_x = allocator_tensor((num_experts_per_rank, capacity, hidden), torch.float8_e4m3fn, allocator=allocator) recv_x_sf = allocator_tensor( From c4928cd50f7b53d925f4c8898383255d6ce438e2 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Tue, 11 Aug 2026 14:59:10 +0800 Subject: [PATCH 04/30] perf(distributed): streamline mega MoE dispatch --- .../mega_moe/example_sm90_fp8_mega_moe.py | 63 ++++++++++--------- 1 file changed, 33 insertions(+), 30 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 51207fa713..4137c9fc68 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -166,30 +166,20 @@ def main(barrier: T.Tensor((num_ranks,), T.int32)): def finalize_routes_kernel( num_tokens: int, - hidden: int, num_experts: int, num_topk: int, num_ranks: int, - capacity: int, threads: int = 128, ): num_experts_per_rank = num_experts // num_ranks - num_scale_groups = hidden // SCALE_GRANULARITY num_routes = num_tokens * num_topk num_work_items = max(num_routes, num_experts_per_rank) @T.prim_func def main( - x_sf: T.Tensor((num_tokens, num_scale_groups), T.float32), topk_idx: T.Tensor((num_tokens, num_topk), T.int32), - topk_weights: T.Tensor((num_tokens, num_topk), T.float32), route_counts: T.Tensor((num_ranks, num_experts), T.int32), recv_counts: T.Tensor((num_experts_per_rank,), T.int32), - recv_x_sf: T.Tensor((num_experts_per_rank, capacity, num_scale_groups), T.float32), - recv_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), - src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), - src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), - src_topk: T.Tensor((num_experts_per_rank, capacity), T.int32), route_slots: T.Tensor((num_tokens, num_topk), T.int32), ): with T.Kernel(T.ceildiv(num_work_items, threads), threads=threads) as bx: @@ -214,19 +204,6 @@ def main( if peer_rank < src_rank[0]: slot += route_counts[peer_rank, expert_idx] route_slots[token_idx, topk_slot] = slot - if slot < capacity: - dst_rank = expert_idx // num_experts_per_rank - local_expert = expert_idx % num_experts_per_rank - T.st(recv_weights[local_expert, slot], topk_weights[token_idx, topk_slot], dst_pe=dst_rank) - T.st(src_ranks[local_expert, slot], src_rank[0], dst_pe=dst_rank) - T.st(src_tokens[local_expert, slot], token_idx, dst_pe=dst_rank) - T.st(src_topk[local_expert, slot], topk_slot, dst_pe=dst_rank) - for scale_idx in T.serial(num_scale_groups): - T.st( - recv_x_sf[local_expert, slot, scale_idx], - x_sf[token_idx, scale_idx], - dst_pe=dst_rank, - ) return main @@ -242,15 +219,26 @@ def dispatch_tokens_kernel( threads: int = 128, ): num_experts_per_rank = num_experts // num_ranks + num_scale_groups = hidden // SCALE_GRANULARITY @T.prim_func def main( x: T.Tensor((num_tokens, hidden), T.float8_e4m3fn), + x_sf: T.Tensor((num_tokens, num_scale_groups), T.float32), topk_idx: T.Tensor((num_tokens, num_topk), T.int32), + topk_weights: T.Tensor((num_tokens, num_topk), T.float32), route_slots: T.Tensor((num_tokens, num_topk), T.int32), recv_x: T.Tensor((num_experts_per_rank, capacity, hidden), T.float8_e4m3fn), + recv_x_sf: T.Tensor((num_experts_per_rank, capacity, num_scale_groups), T.float32), + recv_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), + src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), + src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), + src_topk: T.Tensor((num_experts_per_rank, capacity), T.int32), ): with T.Kernel(T.ceildiv(hidden, block_h), num_tokens * num_topk, threads=threads) as (bx, by): + src_rank = T.alloc_local((1,), T.int32) + src_rank[0] = T.get_rank() + tx = T.get_thread_binding() token_idx = by // num_topk topk_slot = by % num_topk expert_idx = topk_idx[token_idx, topk_slot] @@ -264,7 +252,18 @@ def main( dst_pe=dst_rank, disable_tma=True, ) - T.fence_sys() + if bx == 0 and num_scale_groups > 0: + T.copy( + x_sf[token_idx, :], + recv_x_sf[local_expert, slot, :], + dst_pe=dst_rank, + disable_tma=True, + ) + if tx == 0 and src_rank[0] < num_ranks: + T.st(recv_weights[local_expert, slot], topk_weights[token_idx, topk_slot], dst_pe=dst_rank) + T.st(src_ranks[local_expert, slot], src_rank[0], dst_pe=dst_rank) + T.st(src_tokens[local_expert, slot], token_idx, dst_pe=dst_rank) + T.st(src_topk[local_expert, slot], topk_slot, dst_pe=dst_rank) return main @@ -408,7 +407,6 @@ def main( dst_pe=dst_rank, disable_tma=True, ) - T.fence_sys() return main @@ -571,7 +569,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): assign_local_routes_kernel(num_tokens, num_experts, num_topk, num_ranks), publish_route_counts_kernel(num_experts, num_ranks), device_barrier_kernel(num_ranks), - finalize_routes_kernel(num_tokens, hidden, num_experts, num_topk, num_ranks, capacity), + finalize_routes_kernel(num_tokens, num_experts, num_topk, num_ranks), dispatch_tokens_kernel(num_tokens, hidden, num_experts, num_topk, num_ranks, capacity), fp8_grouped_gemm_kernel(num_experts_per_rank, capacity, 2 * intermediate_hidden, hidden), swiglu_quant_kernel( @@ -684,19 +682,24 @@ def run_pipeline(check_capacity: bool = False): publish_route_counts(route_counts) device_barrier(barrier) finalize_routes( - x_sf, topk_idx, - topk_weights, route_counts, recv_counts, + route_slots, + ) + dispatch_tokens( + x, + x_sf, + topk_idx, + topk_weights, + route_slots, + recv_x, recv_x_sf, recv_weights, src_ranks, src_tokens, src_topk, - route_slots, ) - dispatch_tokens(x, topk_idx, route_slots, recv_x) device_barrier(barrier) if check_capacity: local_max = recv_counts.max() From e50c55da115dd9711640d091a7fcfb1993935203 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Wed, 12 Aug 2026 10:36:36 +0800 Subject: [PATCH 05/30] perf(distributed): fuse mega MoE execution phases --- .../mega_moe/example_sm90_fp8_mega_moe.py | 888 +++++++++++++++--- 1 file changed, 775 insertions(+), 113 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 4137c9fc68..66c2144c8a 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -1,14 +1,14 @@ """Multi-GPU FP8 Mega MoE for SM90 using TileScale distributed primitives. -This implementation uses symmetric VMM buffers for expert dispatch and combine. -The two FP8 GEMMs use per-token/per-128 activation scales and per-(128, 128) -weight scales. Device-side system barriers order the communication phases. +The Flash configuration uses two persistent kernels: one for routing, dispatch, +L1 GEMM, and SwiGLU quantization, and one for L2 GEMM, scatter, and reduction. +The Pro configuration retains the multi-kernel path until its four-warpgroup L1 +kernel is ready. """ from __future__ import annotations import argparse -import math import os from typing import Tuple @@ -67,6 +67,14 @@ def block_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: return x_fp8.view(groups, n, k).contiguous(), scale.contiguous() +def interleave_gate_up_weights(weight: torch.Tensor, granularity: int = 8) -> torch.Tensor: + groups, n, k = weight.shape + half = n // 2 + gate = weight[:, :half].view(groups, half // granularity, granularity, k) + up = weight[:, half:].view(groups, half // granularity, granularity, k) + return torch.stack((gate, up), dim=2).reshape(groups, n, k).contiguous() + + def dequantize_per_token(x: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: m, k = x.shape return (x.float().view(m, k // SCALE_GRANULARITY, SCALE_GRANULARITY) * scale.unsqueeze(-1)).view(m, k) @@ -229,7 +237,10 @@ def main( topk_weights: T.Tensor((num_tokens, num_topk), T.float32), route_slots: T.Tensor((num_tokens, num_topk), T.int32), recv_x: T.Tensor((num_experts_per_rank, capacity, hidden), T.float8_e4m3fn), - recv_x_sf: T.Tensor((num_experts_per_rank, capacity, num_scale_groups), T.float32), + recv_x_sf: T.Tensor( + (num_experts_per_rank, capacity, num_scale_groups), + T.float32, + ), recv_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), @@ -268,6 +279,344 @@ def main( return main +def fused_l1_swiglu_manual_warp_kernel( + num_tokens: int, + hidden: int, + l1_n: int, + num_experts: int, + num_topk: int, + num_ranks: int, + capacity: int, + num_sms: int, + activation_clamp: float = 10.0, + block_h: int = 256, + block_m: int = 64, + block_n: int = 256, + block_k: int = 128, + threads: int = 384, + pipeline_stages: int = 3, +): + num_experts_per_rank = num_experts // num_ranks + num_scale_groups = hidden // SCALE_GRANULARITY + num_routes = num_tokens * num_topk + num_hidden_blocks = ceil_div(hidden, block_h) + num_m_blocks = ceil_div(capacity, block_m) + num_n_blocks = ceil_div(l1_n, block_n) + num_compute_tiles = num_experts_per_rank * num_m_blocks * num_n_blocks + num_k_blocks = hidden // block_k + dispatch_thread = 0 + + @T.prim_func + def main( + x: T.Tensor((num_tokens, hidden), T.float8_e4m3fn), + x_sf: T.Tensor((num_tokens, num_scale_groups), T.float32), + topk_idx: T.Tensor((num_tokens, num_topk), T.int32), + topk_weights: T.Tensor((num_tokens, num_topk), T.float32), + route_counts: T.Tensor((num_ranks, num_experts), T.int32), + recv_counts: T.Tensor((num_experts_per_rank,), T.int32), + route_slots: T.Tensor((num_tokens, num_topk), T.int32), + recv_x: T.Tensor((num_experts_per_rank, capacity, hidden), T.float8_e4m3fn), + recv_x_sf: T.Tensor((num_experts_per_rank, capacity, num_scale_groups), T.float32), + recv_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), + src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), + src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), + src_topk: T.Tensor((num_experts_per_rank, capacity), T.int32), + l1_weight: T.Tensor((num_experts_per_rank, l1_n, hidden), T.float8_e4m3fn), + l1_weight_sf: T.Tensor( + (num_experts_per_rank, l1_n // SCALE_GRANULARITY, hidden // SCALE_GRANULARITY), + T.float32, + ), + l2_x: T.Tensor((num_experts_per_rank, capacity, l1_n // 2), T.float8_e4m3fn), + l2_x_sf: T.Tensor( + (num_experts_per_rank, capacity, l1_n // (2 * SCALE_GRANULARITY)), + T.float32, + ), + barrier: T.Tensor((num_ranks,), T.int32), + ): + with T.Kernel(num_sms, threads=threads) as bid: + tx = T.get_thread_binding() + src_rank = T.alloc_local((1,), T.int32) + src_rank[0] = T.get_rank() + + a_shared = T.alloc_shared( + (pipeline_stages, block_m, block_k), + T.float8_e4m3fn, + ) + b_shared = T.alloc_shared( + (pipeline_stages, block_n, block_k), + T.float8_e4m3fn, + ) + out_shared = T.alloc_shared((block_m, block_n // 2), T.float8_e4m3fn) + stage_barriers = T.alloc_barrier([64] * pipeline_stages + [256] * pipeline_stages) + + if bid == 0: + if tx < 64: + for reset_wave in T.serial(ceil_div(num_experts, 64)): + reset_expert = tx + reset_wave * 64 + if reset_expert < num_experts: + route_counts[src_rank[0], reset_expert] = 0 + T.sync_threads(7, 64) + + for assign_wave in T.serial(ceil_div(num_routes, 64)): + assign_route = tx + assign_wave * 64 + if assign_route < num_routes: + assign_token = assign_route // num_topk + assign_topk = assign_route % num_topk + assign_expert = topk_idx[assign_token, assign_topk] + if assign_expert >= 0 and assign_expert < num_experts: + route_slots[assign_token, assign_topk] = T.atomic_add( + route_counts[src_rank[0], assign_expert], + 1, + memory_order="relaxed", + return_prev=True, + ) + else: + route_slots[assign_token, assign_topk] = -1 + T.sync_threads(7, 64) + + for publish_wave in T.serial(ceil_div(num_experts * num_ranks, 64)): + publish_idx = tx + publish_wave * 64 + if publish_idx < num_experts * num_ranks: + publish_rank = publish_idx // num_experts + publish_expert = publish_idx % num_experts + if publish_rank != src_rank[0]: + T.st( + route_counts[src_rank[0], publish_expert], + route_counts[src_rank[0], publish_expert], + dst_pe=publish_rank, + ) + + T.barrier_blocks(barrier[0]) + + if tx < 64: + if tx < num_experts_per_rank: + recv_count = T.alloc_var(T.int32, init=0) + recv_expert = src_rank[0] * num_experts_per_rank + tx + for count_rank in T.serial(num_ranks): + recv_count += route_counts[count_rank, recv_expert] + recv_counts[tx] = recv_count + + for prefix_wave in T.serial(ceil_div(num_routes, 64)): + prefix_route = tx + prefix_wave * 64 + if prefix_route < num_routes: + prefix_token = prefix_route // num_topk + prefix_topk = prefix_route % num_topk + prefix_expert = topk_idx[prefix_token, prefix_topk] + prefix_slot = T.alloc_var( + T.int32, + init=route_slots[prefix_token, prefix_topk], + ) + if prefix_token < num_tokens and prefix_expert >= 0 and prefix_expert < num_experts and prefix_slot >= 0: + for prefix_rank in T.serial(num_ranks): + if prefix_rank < src_rank[0]: + prefix_slot += route_counts[prefix_rank, prefix_expert] + route_slots[prefix_token, prefix_topk] = prefix_slot + + T.sync_grid() + + if tx < 128: + T.dec_max_nreg(48) + else: + T.inc_max_nreg(208) + + if tx < 64: + for dispatch_wave in T.serial(ceil_div(num_routes, num_sms)): + dispatch_route = bid + dispatch_wave * num_sms + if dispatch_route < num_routes: + dispatch_token = dispatch_route // num_topk + dispatch_topk = dispatch_route % num_topk + dispatch_expert = topk_idx[dispatch_token, dispatch_topk] + dispatch_slot = route_slots[dispatch_token, dispatch_topk] + if dispatch_expert >= 0 and dispatch_slot >= 0 and dispatch_slot < capacity and num_scale_groups > 0 and hidden > 0: + dispatch_rank = dispatch_expert // num_experts_per_rank + dispatch_local_expert = dispatch_expert % num_experts_per_rank + for dispatch_h in T.serial(num_hidden_blocks): + T.copy( + x[ + dispatch_token, + dispatch_h * block_h : (dispatch_h + 1) * block_h, + ], + recv_x[ + dispatch_local_expert, + dispatch_slot, + dispatch_h * block_h : (dispatch_h + 1) * block_h, + ], + dst_pe=dispatch_rank, + disable_tma=True, + ) + T.copy( + x_sf[dispatch_token, :], + recv_x_sf[dispatch_local_expert, dispatch_slot, :], + dst_pe=dispatch_rank, + disable_tma=True, + ) + if tx == dispatch_thread: + T.st( + recv_weights[dispatch_local_expert, dispatch_slot], + topk_weights[dispatch_token, dispatch_topk], + dst_pe=dispatch_rank, + ) + T.st( + src_ranks[dispatch_local_expert, dispatch_slot], + src_rank[0], + dst_pe=dispatch_rank, + ) + T.st( + src_tokens[dispatch_local_expert, dispatch_slot], + dispatch_token, + dst_pe=dispatch_rank, + ) + T.st( + src_topk[dispatch_local_expert, dispatch_slot], + dispatch_topk, + dst_pe=dispatch_rank, + ) + T.fence_sys() + T.sync_grid() + if bid == 0: + T.barrier_blocks(barrier[0]) + T.sync_grid() + + if tx >= 64 and tx < 128: + producer_step = T.alloc_var(T.int32, init=0) + for producer_wave in T.serial(ceil_div(num_compute_tiles, num_sms)): + producer_tile = bid + producer_wave * num_sms + if producer_tile < num_compute_tiles: + producer_n = producer_tile % num_n_blocks + producer_m = (producer_tile // num_n_blocks) % num_m_blocks + producer_expert = producer_tile // (num_n_blocks * num_m_blocks) + if producer_n * block_n < l1_n and producer_m * block_m < recv_counts[producer_expert]: + for producer_k in T.serial(num_k_blocks): + producer_stage = (producer_step + producer_k) % pipeline_stages + producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 + T.mbarrier_wait_parity( + stage_barriers[pipeline_stages + producer_stage], + producer_phase ^ 1, + ) + T.tma_copy( + recv_x[ + producer_expert, + producer_m * block_m : (producer_m + 1) * block_m, + producer_k * block_k : (producer_k + 1) * block_k, + ], + a_shared[producer_stage, :, :], + barrier=stage_barriers[producer_stage], + ) + T.tma_copy( + l1_weight[ + producer_expert, + producer_n * block_n : (producer_n + 1) * block_n, + producer_k * block_k : (producer_k + 1) * block_k, + ], + b_shared[producer_stage, :, :], + barrier=stage_barriers[producer_stage], + ) + T.mbarrier_arrive(stage_barriers[producer_stage]) + producer_step += num_k_blocks + + elif tx >= 128: + partial = T.alloc_fragment((block_m, block_n), T.float32) + accum = T.alloc_fragment((block_m, block_n), T.bfloat16) + gate = T.alloc_fragment((block_m, block_n // 2), T.float32) + up = T.alloc_fragment((block_m, block_n // 2), T.float32) + amax = T.alloc_fragment((block_m,), T.float32) + scale = T.alloc_fragment((block_m,), T.float32) + quant_fp8 = T.alloc_fragment((block_m, block_n // 2), T.float8_e4m3fn) + act_scale = T.alloc_fragment((block_m,), T.float32) + weight_scale = T.alloc_local((2,), T.float32) + consumer_step = T.alloc_var(T.int32, init=0) + + for consumer_wave in T.serial(ceil_div(num_compute_tiles, num_sms)): + consumer_tile = bid + consumer_wave * num_sms + if consumer_tile < num_compute_tiles: + consumer_n = consumer_tile % num_n_blocks + consumer_m = (consumer_tile // num_n_blocks) % num_m_blocks + consumer_expert = consumer_tile // (num_n_blocks * num_m_blocks) + if consumer_n * block_n < l1_n and consumer_m * block_m < recv_counts[consumer_expert]: + T.clear(partial) + T.clear(accum) + for consumer_k in T.serial(num_k_blocks): + consumer_stage = (consumer_step + consumer_k) % pipeline_stages + consumer_phase = ((consumer_step + consumer_k) // pipeline_stages) & 1 + T.mbarrier_wait_parity( + stage_barriers[consumer_stage], + consumer_phase, + ) + T.gemm( + a_shared[consumer_stage, :, :], + b_shared[consumer_stage, :, :], + partial, + transpose_B=True, + ) + for i in T.Parallel(block_m): + act_scale[i] = recv_x_sf[ + consumer_expert, + consumer_m * block_m + i, + consumer_k, + ] + weight_scale[0] = l1_weight_sf[consumer_expert, consumer_n, consumer_k] + weight_scale[1] = l1_weight_sf[ + consumer_expert, + consumer_n + num_n_blocks, + consumer_k, + ] + for i, j in T.Parallel(block_m, block_n): + accum[i, j] = ( + T.cast(partial[i, j], T.bfloat16) + * T.cast( + act_scale[i] * weight_scale[(j % 16) // 8], + T.bfloat16, + ) + + accum[i, j] + ) + T.clear(partial) + T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) + consumer_step += num_k_blocks + for i, j in T.Parallel(block_m, block_n // 2): + gate[i, j] = accum[i, (j // 8) * 16 + j % 8] + for i, j in T.Parallel(block_m, block_n // 2): + up[i, j] = accum[i, (j // 8) * 16 + j % 8 + 8] + for i, j in T.Parallel(block_m, block_n // 2): + gate[i, j] = ( + T.min(gate[i, j], activation_clamp) + * T.sigmoid(T.min(gate[i, j], activation_clamp)) + * T.max( + T.min(up[i, j], activation_clamp), + -activation_clamp, + ) + * recv_weights[ + consumer_expert, + consumer_m * block_m + i, + ] + ) + T.reduce_absmax(gate, amax, dim=1) + for i in T.Parallel(block_m): + scale[i] = T.max(amax[i], 1e-4) / FP8_MAX + l2_x_sf[ + consumer_expert, + consumer_m * block_m + i, + consumer_n, + ] = scale[i] + for i, j in T.Parallel(block_m, block_n // 2): + gate[i, j] = T.clamp( + gate[i, j] / scale[i], + -FP8_MAX, + FP8_MAX, + ) + T.copy(gate, quant_fp8) + T.copy(quant_fp8, out_shared) + T.copy( + out_shared, + l2_x[ + consumer_expert, + consumer_m * block_m, + consumer_n * (block_n // 2), + ], + ) + + return main + + def fp8_grouped_gemm_kernel( num_experts_per_rank: int, capacity: int, @@ -319,6 +668,250 @@ def main( return main +def fused_l2_scatter_reduce_manual_warp_kernel( + num_tokens: int, + hidden: int, + intermediate_hidden: int, + num_experts_per_rank: int, + num_topk: int, + num_ranks: int, + capacity: int, + num_sms: int, + block_m: int = 64, + block_n: int = 256, + block_k: int = 128, + reduce_block_m: int = 8, + reduce_block_h: int = 128, + threads: int = 384, + pipeline_stages: int = 3, +): + num_m_blocks = ceil_div(capacity, block_m) + num_n_blocks = ceil_div(hidden, block_n) + num_compute_tiles = num_experts_per_rank * num_m_blocks * num_n_blocks + num_k_blocks = intermediate_hidden // block_k + num_reduce_n_blocks = ceil_div(hidden, reduce_block_h) + num_reduce_m_blocks = ceil_div(num_tokens, reduce_block_m) + num_reduce_tiles = num_reduce_n_blocks * num_reduce_m_blocks + + @T.prim_func + def main( + a: T.Tensor( + (num_experts_per_rank, capacity, intermediate_hidden), + T.float8_e4m3fn, + ), + b: T.Tensor( + (num_experts_per_rank, hidden, intermediate_hidden), + T.float8_e4m3fn, + ), + a_sf: T.Tensor( + ( + num_experts_per_rank, + capacity, + intermediate_hidden // SCALE_GRANULARITY, + ), + T.float32, + ), + b_sf: T.Tensor( + ( + num_experts_per_rank, + hidden // SCALE_GRANULARITY, + intermediate_hidden // SCALE_GRANULARITY, + ), + T.float32, + ), + recv_counts: T.Tensor((num_experts_per_rank,), T.int32), + src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), + src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), + src_topk: T.Tensor((num_experts_per_rank, capacity), T.int32), + combine: T.Tensor((num_tokens, num_topk, hidden), T.bfloat16), + barrier: T.Tensor((num_ranks,), T.int32), + out: T.Tensor((num_tokens, hidden), T.bfloat16), + ): + with T.Kernel(num_sms, threads=threads) as bid: + tx = T.get_thread_binding() + a_shared = T.alloc_shared( + (pipeline_stages, block_m, block_k), + T.float8_e4m3fn, + ) + b_shared = T.alloc_shared( + (pipeline_stages, block_n, block_k), + T.float8_e4m3fn, + ) + out_shared = T.alloc_shared((block_m, block_n), T.bfloat16) + reduce_shared = T.alloc_shared((reduce_block_m, reduce_block_h), T.bfloat16) + stage_barriers = T.alloc_barrier([64] * pipeline_stages + [256] * pipeline_stages) + + if tx < 128: + T.dec_max_nreg(48) + else: + T.inc_max_nreg(208) + + if tx >= 64 and tx < 128: + producer_step = T.alloc_var(T.int32, init=0) + for producer_wave in T.serial(ceil_div(num_compute_tiles, num_sms)): + producer_tile = bid + producer_wave * num_sms + if producer_tile < num_compute_tiles: + producer_n = producer_tile % num_n_blocks + producer_m = (producer_tile // num_n_blocks) % num_m_blocks + producer_expert = producer_tile // (num_n_blocks * num_m_blocks) + if ( + producer_expert < num_experts_per_rank + and producer_n * block_n < hidden + and producer_m * block_m < capacity + and producer_m * block_m < recv_counts[producer_expert] + and num_k_blocks * block_k == intermediate_hidden + ): + for producer_k in T.serial(num_k_blocks): + producer_stage = (producer_step + producer_k) % pipeline_stages + producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 + T.mbarrier_wait_parity( + stage_barriers[pipeline_stages + producer_stage], + producer_phase ^ 1, + ) + T.tma_copy( + a[ + producer_expert, + producer_m * block_m : (producer_m + 1) * block_m, + producer_k * block_k : (producer_k + 1) * block_k, + ], + a_shared[producer_stage, :, :], + barrier=stage_barriers[producer_stage], + ) + T.tma_copy( + b[ + producer_expert, + producer_n * block_n : (producer_n + 1) * block_n, + producer_k * block_k : (producer_k + 1) * block_k, + ], + b_shared[producer_stage, :, :], + barrier=stage_barriers[producer_stage], + ) + T.mbarrier_arrive(stage_barriers[producer_stage]) + producer_step += num_k_blocks + + elif tx >= 128: + partial = T.alloc_fragment((block_m, block_n), T.float32) + accum = T.alloc_fragment((block_m, block_n), T.bfloat16) + act_scale = T.alloc_fragment((block_m,), T.float32) + weight_scale = T.alloc_local((2,), T.float32) + consumer_step = T.alloc_var(T.int32, init=0) + + for consumer_wave in T.serial(ceil_div(num_compute_tiles, num_sms)): + consumer_tile = bid + consumer_wave * num_sms + if consumer_tile < num_compute_tiles: + consumer_n = consumer_tile % num_n_blocks + consumer_m = (consumer_tile // num_n_blocks) % num_m_blocks + consumer_expert = consumer_tile // (num_n_blocks * num_m_blocks) + if ( + consumer_expert < num_experts_per_rank + and consumer_n * block_n < hidden + and consumer_m * block_m < capacity + and consumer_m * block_m < recv_counts[consumer_expert] + and num_k_blocks * block_k == intermediate_hidden + ): + T.clear(partial) + T.clear(accum) + for consumer_k in T.serial(num_k_blocks): + consumer_stage = (consumer_step + consumer_k) % pipeline_stages + consumer_phase = ((consumer_step + consumer_k) // pipeline_stages) & 1 + T.mbarrier_wait_parity( + stage_barriers[consumer_stage], + consumer_phase, + ) + T.gemm( + a_shared[consumer_stage, :, :], + b_shared[consumer_stage, :, :], + partial, + transpose_B=True, + ) + for i in T.Parallel(block_m): + act_scale[i] = a_sf[ + consumer_expert, + consumer_m * block_m + i, + consumer_k, + ] + weight_scale[0] = b_sf[ + consumer_expert, + consumer_n * 2, + consumer_k, + ] + weight_scale[1] = b_sf[ + consumer_expert, + consumer_n * 2 + 1, + consumer_k, + ] + for i, j in T.Parallel(block_m, block_n): + accum[i, j] = ( + T.cast(partial[i, j], T.bfloat16) + * T.cast( + act_scale[i] * weight_scale[j // 128], + T.bfloat16, + ) + + accum[i, j] + ) + T.clear(partial) + T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) + consumer_step += num_k_blocks + T.copy(accum, out_shared) + for row in T.serial(block_m): + pool_row = consumer_m * block_m + row + if pool_row < recv_counts[consumer_expert]: + dst_rank = src_ranks[consumer_expert, pool_row] + dst_token = src_tokens[consumer_expert, pool_row] + dst_topk = src_topk[consumer_expert, pool_row] + if ( + dst_rank >= 0 + and dst_rank < num_ranks + and dst_token >= 0 + and dst_token < num_tokens + and dst_topk >= 0 + and dst_topk < num_topk + ): + T.copy( + out_shared[row, :], + combine[ + dst_token, + dst_topk, + consumer_n * block_n : (consumer_n + 1) * block_n, + ], + dst_pe=dst_rank, + disable_tma=True, + ) + + T.fence_sys() + T.sync_grid() + if bid == 0: + T.barrier_blocks(barrier[0]) + T.sync_grid() + + if tx < 128: + reduce_accum = T.alloc_fragment((reduce_block_m, reduce_block_h), T.float32) + for reduce_wave in T.serial(ceil_div(num_reduce_tiles, num_sms)): + reduce_tile = bid + reduce_wave * num_sms + if reduce_tile < num_reduce_tiles: + reduce_n = reduce_tile % num_reduce_n_blocks + reduce_m = reduce_tile // num_reduce_n_blocks + T.clear(reduce_accum) + for topk_slot in T.serial(num_topk): + for i, j in T.Parallel(reduce_block_m, reduce_block_h): + if reduce_m * reduce_block_m + i < num_tokens: + reduce_accum[i, j] += combine[ + reduce_m * reduce_block_m + i, + topk_slot, + reduce_n * reduce_block_h + j, + ] + T.copy(reduce_accum, reduce_shared) + T.copy( + reduce_shared, + out[ + reduce_m * reduce_block_m, + reduce_n * reduce_block_h, + ], + ) + + return main + + def swiglu_quant_kernel( num_experts_per_rank: int, capacity: int, @@ -359,12 +952,7 @@ def main( for i, j in T.Parallel(block_m, block_n): gate[i, j] = T.min(gate[i, j], activation_clamp) up[i, j] = T.max(T.min(up[i, j], activation_clamp), -activation_clamp) - activated[i, j] = ( - gate[i, j] - * T.sigmoid(gate[i, j]) - * up[i, j] - * route_weights[bz, by * block_m + i] - ) + activated[i, j] = gate[i, j] * T.sigmoid(gate[i, j]) * up[i, j] * route_weights[bz, by * block_m + i] T.reduce_absmax(activated, amax, dim=1) for i in T.Parallel(block_m): scale[i] = T.max(amax[i], 1e-4) / FP8_MAX @@ -450,25 +1038,24 @@ def _allocator_size_bytes( bf16 = 2 fp32 = 4 i32 = 4 - weight_bytes = num_experts_per_rank * ( - 2 * intermediate_hidden * hidden * fp8 + hidden * intermediate_hidden * fp8 - ) - weight_scale_bytes = num_experts_per_rank * ( - (2 * intermediate_hidden // 128) * (hidden // 128) - + (hidden // 128) * (intermediate_hidden // 128) - ) * fp32 - pool_bytes = num_experts_per_rank * capacity * ( - hidden * fp8 - + (hidden // 128) * fp32 - + 4 * i32 - + 2 * intermediate_hidden * bf16 - + intermediate_hidden * fp8 - + (intermediate_hidden // 128) * fp32 - + hidden * bf16 + weight_bytes = num_experts_per_rank * (2 * intermediate_hidden * hidden * fp8 + hidden * intermediate_hidden * fp8) + weight_scale_bytes = ( + num_experts_per_rank * ((2 * intermediate_hidden // 128) * (hidden // 128) + (hidden // 128) * (intermediate_hidden // 128)) * fp32 ) - input_bytes = num_tokens * ( - hidden * fp8 + (hidden // 128) * fp32 + num_topk * (3 * i32 + fp32) + num_topk * hidden * bf16 + pool_bytes = ( + num_experts_per_rank + * capacity + * ( + hidden * fp8 + + (hidden // 128) * fp32 + + 4 * i32 + + 2 * intermediate_hidden * bf16 + + intermediate_hidden * fp8 + + (intermediate_hidden // 128) * fp32 + + hidden * bf16 + ) ) + input_bytes = num_tokens * (hidden * fp8 + (hidden // 128) * fp32 + num_topk * (3 * i32 + fp32) + num_topk * hidden * bf16) return align_up(weight_bytes + weight_scale_bytes + pool_bytes + input_bytes + 2**27, 2**20) @@ -538,6 +1125,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): num_topk = model["num_topk"] num_tokens = args.num_tokens activation_clamp = args.activation_clamp + use_fused = args.model_config != "pro" assert num_experts % num_local_ranks == 0 assert hidden % 256 == 0 and intermediate_hidden % 128 == 0 @@ -547,6 +1135,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): rank, num_ranks, group = init_dist(local_rank, num_local_ranks) assert rank == local_rank and num_ranks == num_local_ranks + num_sms = torch.cuda.get_device_properties(local_rank).multi_processor_count allocator = get_allocator( size=_allocator_size_bytes( num_tokens, @@ -564,40 +1153,68 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): use_vmm=True, ) - kernel_specs = [ - reset_route_counts_kernel(num_experts, num_ranks), - assign_local_routes_kernel(num_tokens, num_experts, num_topk, num_ranks), - publish_route_counts_kernel(num_experts, num_ranks), - device_barrier_kernel(num_ranks), - finalize_routes_kernel(num_tokens, num_experts, num_topk, num_ranks), - dispatch_tokens_kernel(num_tokens, hidden, num_experts, num_topk, num_ranks, capacity), - fp8_grouped_gemm_kernel(num_experts_per_rank, capacity, 2 * intermediate_hidden, hidden), - swiglu_quant_kernel( - num_experts_per_rank, - capacity, - intermediate_hidden, - activation_clamp=activation_clamp, - ), - fp8_grouped_gemm_kernel(num_experts_per_rank, capacity, hidden, intermediate_hidden), - scatter_outputs_kernel(num_experts_per_rank, capacity, num_tokens, num_topk, hidden), - reduce_topk_kernel(num_tokens, num_topk, hidden), - ] + if use_fused: + kernel_specs = [ + fused_l1_swiglu_manual_warp_kernel( + num_tokens, + hidden, + 2 * intermediate_hidden, + num_experts, + num_topk, + num_ranks, + capacity, + num_sms, + activation_clamp=activation_clamp, + ), + fused_l2_scatter_reduce_manual_warp_kernel( + num_tokens, + hidden, + intermediate_hidden, + num_experts_per_rank, + num_topk, + num_ranks, + capacity, + num_sms, + ), + ] + else: + kernel_specs = [ + reset_route_counts_kernel(num_experts, num_ranks), + assign_local_routes_kernel(num_tokens, num_experts, num_topk, num_ranks), + publish_route_counts_kernel(num_experts, num_ranks), + device_barrier_kernel(num_ranks), + finalize_routes_kernel(num_tokens, num_experts, num_topk, num_ranks), + dispatch_tokens_kernel(num_tokens, hidden, num_experts, num_topk, num_ranks, capacity), + fp8_grouped_gemm_kernel(num_experts_per_rank, capacity, 2 * intermediate_hidden, hidden), + swiglu_quant_kernel( + num_experts_per_rank, + capacity, + intermediate_hidden, + activation_clamp=activation_clamp, + ), + fp8_grouped_gemm_kernel(num_experts_per_rank, capacity, hidden, intermediate_hidden), + scatter_outputs_kernel(num_experts_per_rank, capacity, num_tokens, num_topk, hidden), + reduce_topk_kernel(num_tokens, num_topk, hidden), + ] kernels = [tilelang.compile(spec, compile_once=True, compile_group=group) for spec in kernel_specs] for kernel in kernels: kernel.initialize(allocator=allocator) - ( - reset_route_counts, - assign_local_routes, - publish_route_counts, - device_barrier, - finalize_routes, - dispatch_tokens, - l1_gemm, - swiglu_quant, - l2_gemm, - scatter_outputs, - reduce_topk, - ) = kernels + if use_fused: + fused_l1, fused_l2 = kernels + else: + ( + reset_route_counts, + assign_local_routes, + publish_route_counts, + device_barrier, + finalize_routes, + dispatch_tokens, + l1_gemm, + swiglu_quant, + l2_gemm, + scatter_outputs, + reduce_topk, + ) = kernels if local_rank == 0 and args.print_source: for kernel in kernels: @@ -610,16 +1227,22 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): topk_weights_src, topk_idx_src = torch.topk(scores, num_topk, dim=-1, sorted=False) topk_idx_src = topk_idx_src.to(torch.int32) - l1_bf16 = torch.randn( - (num_experts_per_rank, 2 * intermediate_hidden, hidden), - dtype=torch.bfloat16, - device="cuda", - ) * 0.05 - l2_bf16 = torch.randn( - (num_experts_per_rank, hidden, intermediate_hidden), - dtype=torch.bfloat16, - device="cuda", - ) * 0.05 + l1_bf16 = ( + torch.randn( + (num_experts_per_rank, 2 * intermediate_hidden, hidden), + dtype=torch.bfloat16, + device="cuda", + ) + * 0.05 + ) + l2_bf16 = ( + torch.randn( + (num_experts_per_rank, hidden, intermediate_hidden), + dtype=torch.bfloat16, + device="cuda", + ) + * 0.05 + ) l1_fp8_src, l1_sf_src = block_cast_to_fp8(l1_bf16) l2_fp8_src, l2_sf_src = block_cast_to_fp8(l2_bf16) del scores, l1_bf16, l2_bf16 @@ -630,10 +1253,10 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): x = allocator_tensor(x_fp8_src.shape, x_fp8_src.dtype, allocator=allocator).copy_(x_fp8_src) x_sf = allocator_tensor(x_sf_src.shape, x_sf_src.dtype, allocator=allocator).copy_(x_sf_src) topk_idx = allocator_tensor(topk_idx_src.shape, topk_idx_src.dtype, allocator=allocator).copy_(topk_idx_src) - topk_weights = allocator_tensor( - topk_weights_src.shape, topk_weights_src.dtype, allocator=allocator - ).copy_(topk_weights_src) - l1_fp8 = allocator_tensor(l1_fp8_src.shape, l1_fp8_src.dtype, allocator=allocator).copy_(l1_fp8_src) + topk_weights = allocator_tensor(topk_weights_src.shape, topk_weights_src.dtype, allocator=allocator).copy_(topk_weights_src) + l1_fp8_kernel = interleave_gate_up_weights(l1_fp8_src) if use_fused else l1_fp8_src + l1_fp8 = allocator_tensor(l1_fp8_kernel.shape, l1_fp8_kernel.dtype, allocator=allocator).copy_(l1_fp8_kernel) + del l1_fp8_kernel l1_sf = allocator_tensor(l1_sf_src.shape, l1_sf_src.dtype, allocator=allocator).copy_(l1_sf_src) l2_fp8 = allocator_tensor(l2_fp8_src.shape, l2_fp8_src.dtype, allocator=allocator).copy_(l2_fp8_src) l2_sf = allocator_tensor(l2_sf_src.shape, l2_sf_src.dtype, allocator=allocator).copy_(l2_sf_src) @@ -642,25 +1265,33 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): recv_counts = allocator_tensor((num_experts_per_rank,), torch.int32, allocator=allocator) recv_x = allocator_tensor((num_experts_per_rank, capacity, hidden), torch.float8_e4m3fn, allocator=allocator) recv_x_sf = allocator_tensor( - (num_experts_per_rank, capacity, hidden // SCALE_GRANULARITY), torch.float32, allocator=allocator + (num_experts_per_rank, capacity, hidden // SCALE_GRANULARITY), + torch.float32, + allocator=allocator, ) recv_weights = allocator_tensor((num_experts_per_rank, capacity), torch.float32, allocator=allocator) src_ranks = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) src_tokens = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) src_topk = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) route_slots = allocator_tensor((num_tokens, num_topk), torch.int32, allocator=allocator) - l1_out = allocator_tensor( - (num_experts_per_rank, capacity, 2 * intermediate_hidden), torch.bfloat16, allocator=allocator - ) - l2_x = allocator_tensor( - (num_experts_per_rank, capacity, intermediate_hidden), torch.float8_e4m3fn, allocator=allocator - ) + if not use_fused: + l1_out = allocator_tensor( + (num_experts_per_rank, capacity, 2 * intermediate_hidden), + torch.bfloat16, + allocator=allocator, + ) + l2_x = allocator_tensor((num_experts_per_rank, capacity, intermediate_hidden), torch.float8_e4m3fn, allocator=allocator) l2_x_sf = allocator_tensor( (num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), torch.float32, allocator=allocator, ) - l2_out = allocator_tensor((num_experts_per_rank, capacity, hidden), torch.bfloat16, allocator=allocator) + if not use_fused: + l2_out = allocator_tensor( + (num_experts_per_rank, capacity, hidden), + torch.bfloat16, + allocator=allocator, + ) combine = allocator_tensor((num_tokens, num_topk, hidden), torch.bfloat16, allocator=allocator) out = allocator_tensor((num_tokens, hidden), torch.bfloat16, allocator=allocator) @@ -677,42 +1308,72 @@ def reset_state(): dist.barrier(group=group) def run_pipeline(check_capacity: bool = False): - reset_route_counts(route_counts) - assign_local_routes(topk_idx, route_counts, route_slots) - publish_route_counts(route_counts) - device_barrier(barrier) - finalize_routes( - topk_idx, - route_counts, - recv_counts, - route_slots, - ) - dispatch_tokens( - x, - x_sf, - topk_idx, - topk_weights, - route_slots, - recv_x, - recv_x_sf, - recv_weights, - src_ranks, - src_tokens, - src_topk, - ) - device_barrier(barrier) + if use_fused: + fused_l1( + x, + x_sf, + topk_idx, + topk_weights, + route_counts, + recv_counts, + route_slots, + recv_x, + recv_x_sf, + recv_weights, + src_ranks, + src_tokens, + src_topk, + l1_fp8, + l1_sf, + l2_x, + l2_x_sf, + barrier, + ) + else: + reset_route_counts(route_counts) + assign_local_routes(topk_idx, route_counts, route_slots) + publish_route_counts(route_counts) + device_barrier(barrier) + finalize_routes(topk_idx, route_counts, recv_counts, route_slots) + dispatch_tokens( + x, + x_sf, + topk_idx, + topk_weights, + route_slots, + recv_x, + recv_x_sf, + recv_weights, + src_ranks, + src_tokens, + src_topk, + ) + device_barrier(barrier) if check_capacity: local_max = recv_counts.max() dist.all_reduce(local_max, op=dist.ReduceOp.MAX, group=group) - assert local_max.item() <= capacity, ( - f"expert capacity {capacity} is smaller than received routes {local_max.item()}" + assert local_max.item() <= capacity, f"expert capacity {capacity} is smaller than received routes {local_max.item()}" + if use_fused: + fused_l2( + l2_x, + l2_fp8, + l2_x_sf, + l2_sf, + recv_counts, + src_ranks, + src_tokens, + src_topk, + combine, + barrier, + out, ) - l1_gemm(recv_x, l1_fp8, recv_x_sf, l1_sf, recv_counts, l1_out) - swiglu_quant(l1_out, recv_weights, recv_counts, l2_x, l2_x_sf) - l2_gemm(l2_x, l2_fp8, l2_x_sf, l2_sf, recv_counts, l2_out) - scatter_outputs(l2_out, recv_counts, src_ranks, src_tokens, src_topk, combine) - device_barrier(barrier) - reduce_topk(combine, out) + else: + l1_gemm(recv_x, l1_fp8, recv_x_sf, l1_sf, recv_counts, l1_out) + swiglu_quant(l1_out, recv_weights, recv_counts, l2_x, l2_x_sf) + l2_gemm(l2_x, l2_fp8, l2_x_sf, l2_sf, recv_counts, l2_out) + scatter_outputs(l2_out, recv_counts, src_ranks, src_tokens, src_topk, combine) + device_barrier(barrier) + reduce_topk(combine, out) return out reset_state() @@ -749,6 +1410,7 @@ def run_pipeline(check_capacity: bool = False): if local_rank == 0: print( f"tilescale sm90 fp8 mega moe: model={args.model_config} M={num_tokens} " + f"implementation={'fused' if use_fused else 'multi-kernel'} " f"capacity={capacity} latency={latency * 1000:.1f} us" ) From 8d08b0ef3962c9e5bdf825dcf4395a9320ad839d Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Wed, 12 Aug 2026 11:31:41 +0800 Subject: [PATCH 06/30] perf(distributed): overlap mega MoE dispatch and compute --- .../mega_moe/example_sm90_fp8_mega_moe.py | 83 ++++++++++++------- 1 file changed, 54 insertions(+), 29 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 66c2144c8a..55136c75a9 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -289,7 +289,6 @@ def fused_l1_swiglu_manual_warp_kernel( capacity: int, num_sms: int, activation_clamp: float = 10.0, - block_h: int = 256, block_m: int = 64, block_n: int = 256, block_k: int = 128, @@ -299,7 +298,6 @@ def fused_l1_swiglu_manual_warp_kernel( num_experts_per_rank = num_experts // num_ranks num_scale_groups = hidden // SCALE_GRANULARITY num_routes = num_tokens * num_topk - num_hidden_blocks = ceil_div(hidden, block_h) num_m_blocks = ceil_div(capacity, block_m) num_n_blocks = ceil_div(l1_n, block_n) num_compute_tiles = num_experts_per_rank * num_m_blocks * num_n_blocks @@ -315,6 +313,7 @@ def main( route_counts: T.Tensor((num_ranks, num_experts), T.int32), recv_counts: T.Tensor((num_experts_per_rank,), T.int32), route_slots: T.Tensor((num_tokens, num_topk), T.int32), + arrivals: T.Tensor((num_experts_per_rank, num_m_blocks), T.uint32), recv_x: T.Tensor((num_experts_per_rank, capacity, hidden), T.float8_e4m3fn), recv_x_sf: T.Tensor((num_experts_per_rank, capacity, num_scale_groups), T.float32), recv_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), @@ -420,8 +419,10 @@ def main( T.inc_max_nreg(208) if tx < 64: - for dispatch_wave in T.serial(ceil_div(num_routes, num_sms)): - dispatch_route = bid + dispatch_wave * num_sms + dispatch_warp = tx // 32 + dispatch_lane = tx % 32 + for dispatch_wave in T.serial(ceil_div(num_routes, num_sms * 2)): + dispatch_route = bid * 2 + dispatch_warp + dispatch_wave * num_sms * 2 if dispatch_route < num_routes: dispatch_token = dispatch_route // num_topk dispatch_topk = dispatch_route % num_topk @@ -430,27 +431,19 @@ def main( if dispatch_expert >= 0 and dispatch_slot >= 0 and dispatch_slot < capacity and num_scale_groups > 0 and hidden > 0: dispatch_rank = dispatch_expert // num_experts_per_rank dispatch_local_expert = dispatch_expert % num_experts_per_rank - for dispatch_h in T.serial(num_hidden_blocks): - T.copy( - x[ - dispatch_token, - dispatch_h * block_h : (dispatch_h + 1) * block_h, - ], - recv_x[ - dispatch_local_expert, - dispatch_slot, - dispatch_h * block_h : (dispatch_h + 1) * block_h, - ], - dst_pe=dispatch_rank, - disable_tma=True, - ) - T.copy( - x_sf[dispatch_token, :], - recv_x_sf[dispatch_local_expert, dispatch_slot, :], + T.put_warp( + T.address_of(x[dispatch_token, 0]), + T.address_of(recv_x[dispatch_local_expert, dispatch_slot, 0]), + hidden, + dst_pe=dispatch_rank, + ) + T.put_warp( + T.address_of(x_sf[dispatch_token, 0]), + T.address_of(recv_x_sf[dispatch_local_expert, dispatch_slot, 0]), + num_scale_groups, dst_pe=dispatch_rank, - disable_tma=True, ) - if tx == dispatch_thread: + if dispatch_lane == dispatch_thread: T.st( recv_weights[dispatch_local_expert, dispatch_slot], topk_weights[dispatch_token, dispatch_topk], @@ -471,11 +464,18 @@ def main( dispatch_topk, dst_pe=dispatch_rank, ) - T.fence_sys() - T.sync_grid() - if bid == 0: - T.barrier_blocks(barrier[0]) - T.sync_grid() + T.sync_warp() + if dispatch_lane == dispatch_thread: + T.fence_sys() + T.atomic_add( + arrivals[ + dispatch_local_expert, + dispatch_slot // block_m, + ], + 1, + memory_order="relaxed", + dst_pe=dispatch_rank, + ) if tx >= 64 and tx < 128: producer_step = T.alloc_var(T.int32, init=0) @@ -486,6 +486,18 @@ def main( producer_m = (producer_tile // num_n_blocks) % num_m_blocks producer_expert = producer_tile // (num_n_blocks * num_m_blocks) if producer_n * block_n < l1_n and producer_m * block_m < recv_counts[producer_expert]: + producer_arrivals = T.min( + block_m, + recv_counts[producer_expert] - producer_m * block_m, + ) + if tx == 64: + T.wait_ge( + arrivals[producer_expert, producer_m], + producer_arrivals, + scope=T.WaitScope.SYS, + semantics=T.WaitSemantics.ACQUIRE, + ) + T.sync_threads(5, 64) for producer_k in T.serial(num_k_blocks): producer_stage = (producer_step + producer_k) % pipeline_stages producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 @@ -1274,6 +1286,12 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): src_tokens = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) src_topk = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) route_slots = allocator_tensor((num_tokens, num_topk), torch.int32, allocator=allocator) + if use_fused: + arrivals = allocator_tensor( + (num_experts_per_rank, ceil_div(capacity, 64)), + torch.uint32, + allocator=allocator, + ) if not use_fused: l1_out = allocator_tensor( (num_experts_per_rank, capacity, 2 * intermediate_hidden), @@ -1299,6 +1317,8 @@ def reset_state(): route_counts.zero_() barrier.zero_() recv_counts.zero_() + if use_fused: + arrivals.zero_() recv_x.zero_() recv_x_sf.zero_() recv_weights.zero_() @@ -1317,6 +1337,7 @@ def run_pipeline(check_capacity: bool = False): route_counts, recv_counts, route_slots, + arrivals, recv_x, recv_x_sf, recv_weights, @@ -1400,9 +1421,13 @@ def run_pipeline(check_capacity: bool = False): if args.rep > 0: reset_state() + # Stateful synchronization counters must be reset between warmup iterations. + for _ in range(args.warmup): + run_pipeline() + reset_state() latency = do_bench( run_pipeline, - warmup=args.warmup, + warmup=0, rep=args.rep, post_fn=reset_state, group=group, From bf8195481acb9161d63868b88f563995d54acfe8 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Wed, 12 Aug 2026 14:17:52 +0800 Subject: [PATCH 07/30] perf(distributed): accelerate mega MoE destination dispatch --- .../mega_moe/example_sm90_fp8_mega_moe.py | 130 ++++++++++-------- 1 file changed, 73 insertions(+), 57 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 55136c75a9..5941a183a8 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -293,7 +293,7 @@ def fused_l1_swiglu_manual_warp_kernel( block_n: int = 256, block_k: int = 128, threads: int = 384, - pipeline_stages: int = 3, + pipeline_stages: int = 5, ): num_experts_per_rank = num_experts // num_ranks num_scale_groups = hidden // SCALE_GRANULARITY @@ -418,64 +418,80 @@ def main( else: T.inc_max_nreg(208) + dispatch_warp = tx // 32 + dispatch_lane = tx % 32 if tx < 64: - dispatch_warp = tx // 32 - dispatch_lane = tx % 32 - for dispatch_wave in T.serial(ceil_div(num_routes, num_sms * 2)): - dispatch_route = bid * 2 + dispatch_warp + dispatch_wave * num_sms * 2 - if dispatch_route < num_routes: - dispatch_token = dispatch_route // num_topk - dispatch_topk = dispatch_route % num_topk - dispatch_expert = topk_idx[dispatch_token, dispatch_topk] - dispatch_slot = route_slots[dispatch_token, dispatch_topk] - if dispatch_expert >= 0 and dispatch_slot >= 0 and dispatch_slot < capacity and num_scale_groups > 0 and hidden > 0: - dispatch_rank = dispatch_expert // num_experts_per_rank - dispatch_local_expert = dispatch_expert % num_experts_per_rank - T.put_warp( - T.address_of(x[dispatch_token, 0]), - T.address_of(recv_x[dispatch_local_expert, dispatch_slot, 0]), - hidden, - dst_pe=dispatch_rank, + for metadata_wave in T.serial(ceil_div(num_routes, num_sms * 64)): + metadata_route = bid * 64 + tx + metadata_wave * num_sms * 64 + if metadata_route < num_routes: + metadata_token = metadata_route // num_topk + metadata_topk = metadata_route % num_topk + metadata_expert = topk_idx[metadata_token, metadata_topk] + metadata_slot = route_slots[metadata_token, metadata_topk] + if metadata_expert >= 0 and metadata_slot >= 0 and metadata_slot < capacity: + metadata_rank = metadata_expert // num_experts_per_rank + metadata_local_expert = metadata_expert % num_experts_per_rank + T.st( + recv_weights[metadata_local_expert, metadata_slot], + topk_weights[metadata_token, metadata_topk], + dst_pe=metadata_rank, ) - T.put_warp( - T.address_of(x_sf[dispatch_token, 0]), - T.address_of(recv_x_sf[dispatch_local_expert, dispatch_slot, 0]), - num_scale_groups, - dst_pe=dispatch_rank, + T.st( + src_tokens[metadata_local_expert, metadata_slot], + metadata_token, + dst_pe=metadata_rank, + ) + T.st( + src_topk[metadata_local_expert, metadata_slot], + metadata_topk, + dst_pe=metadata_rank, + ) + T.st( + src_ranks[metadata_local_expert, metadata_slot], + src_rank[0], + scope="sys", + sem="release", + dst_pe=metadata_rank, + ) + + if tx < 64: + for pull_wave in T.serial(ceil_div(num_experts_per_rank * capacity, num_sms * 2)): + pull_idx = bid * 2 + dispatch_warp + pull_wave * num_sms * 2 + pull_expert = pull_idx // capacity + pull_slot = pull_idx % capacity + if pull_expert < num_experts_per_rank and pull_slot < recv_counts[pull_expert]: + if dispatch_lane == dispatch_thread: + T.wait_ge( + src_ranks[pull_expert, pull_slot], + 0, + scope=T.WaitScope.SYS, + semantics=T.WaitSemantics.ACQUIRE, + ) + T.sync_warp() + pull_rank = src_ranks[pull_expert, pull_slot] + pull_token = src_tokens[pull_expert, pull_slot] + T.get_warp( + T.address_of(x[pull_token, 0]), + T.address_of(recv_x[pull_expert, pull_slot, 0]), + hidden, + src_pe=pull_rank, + unroll_factor=8, + ) + T.get_warp( + T.address_of(x_sf[pull_token, 0]), + T.address_of(recv_x_sf[pull_expert, pull_slot, 0]), + num_scale_groups, + src_pe=pull_rank, + unroll_factor=8, + ) + T.sync_warp() + if dispatch_lane == dispatch_thread: + T.atom_add( + arrivals[pull_expert, pull_slot // block_m], + 1, + scope="gpu", + sem="release", ) - if dispatch_lane == dispatch_thread: - T.st( - recv_weights[dispatch_local_expert, dispatch_slot], - topk_weights[dispatch_token, dispatch_topk], - dst_pe=dispatch_rank, - ) - T.st( - src_ranks[dispatch_local_expert, dispatch_slot], - src_rank[0], - dst_pe=dispatch_rank, - ) - T.st( - src_tokens[dispatch_local_expert, dispatch_slot], - dispatch_token, - dst_pe=dispatch_rank, - ) - T.st( - src_topk[dispatch_local_expert, dispatch_slot], - dispatch_topk, - dst_pe=dispatch_rank, - ) - T.sync_warp() - if dispatch_lane == dispatch_thread: - T.fence_sys() - T.atomic_add( - arrivals[ - dispatch_local_expert, - dispatch_slot // block_m, - ], - 1, - memory_order="relaxed", - dst_pe=dispatch_rank, - ) if tx >= 64 and tx < 128: producer_step = T.alloc_var(T.int32, init=0) @@ -494,7 +510,7 @@ def main( T.wait_ge( arrivals[producer_expert, producer_m], producer_arrivals, - scope=T.WaitScope.SYS, + scope=T.WaitScope.GPU, semantics=T.WaitSemantics.ACQUIRE, ) T.sync_threads(5, 64) From fa2f081abef309e25af914d2780b7f4c2d127194 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Wed, 12 Aug 2026 14:42:21 +0800 Subject: [PATCH 08/30] perf(distributed): enable fast math for mega MoE L1 --- examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py | 1 + 1 file changed, 1 insertion(+) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 5941a183a8..b4dd4489c8 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -332,6 +332,7 @@ def main( ), barrier: T.Tensor((num_ranks,), T.int32), ): + T.annotate_pass_configs({tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) with T.Kernel(num_sms, threads=threads) as bid: tx = T.get_thread_binding() src_rank = T.alloc_local((1,), T.int32) From 2c7dc26f4dee87bd2a32ad6415d9deb7872c6204 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Wed, 12 Aug 2026 14:54:21 +0800 Subject: [PATCH 09/30] perf(distributed): vectorize mega MoE output scatter --- .../mega_moe/example_sm90_fp8_mega_moe.py | 23 +++++++++++-------- 1 file changed, 14 insertions(+), 9 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index b4dd4489c8..c191bc9c9c 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -882,7 +882,9 @@ def main( T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) consumer_step += num_k_blocks T.copy(accum, out_shared) - for row in T.serial(block_m): + scatter_warp = (tx - 128) // 32 + for row_in_warp in T.serial(block_m // 8): + row = scatter_warp * (block_m // 8) + row_in_warp pool_row = consumer_m * block_m + row if pool_row < recv_counts[consumer_expert]: dst_rank = src_ranks[consumer_expert, pool_row] @@ -896,15 +898,18 @@ def main( and dst_topk >= 0 and dst_topk < num_topk ): - T.copy( - out_shared[row, :], - combine[ - dst_token, - dst_topk, - consumer_n * block_n : (consumer_n + 1) * block_n, - ], + T.put_warp( + T.address_of(out_shared[row, 0]), + T.address_of( + combine[ + dst_token, + dst_topk, + consumer_n * block_n, + ] + ), + block_n, dst_pe=dst_rank, - disable_tma=True, + unroll_factor=1, ) T.fence_sys() From 2c25e36e4bbc3fcabf1636896dadb2075d414404 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Wed, 12 Aug 2026 16:04:47 +0800 Subject: [PATCH 10/30] perf(distributed): parallelize mega MoE routing --- .../mega_moe/example_sm90_fp8_mega_moe.py | 25 ++++++++++--------- 1 file changed, 13 insertions(+), 12 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index c191bc9c9c..97cfd77fc3 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -303,6 +303,7 @@ def fused_l1_swiglu_manual_warp_kernel( num_compute_tiles = num_experts_per_rank * num_m_blocks * num_n_blocks num_k_blocks = hidden // block_k dispatch_thread = 0 + route_threads = 256 @T.prim_func def main( @@ -350,15 +351,15 @@ def main( stage_barriers = T.alloc_barrier([64] * pipeline_stages + [256] * pipeline_stages) if bid == 0: - if tx < 64: - for reset_wave in T.serial(ceil_div(num_experts, 64)): - reset_expert = tx + reset_wave * 64 + if tx < route_threads: + for reset_wave in T.serial(ceil_div(num_experts, route_threads)): + reset_expert = tx + reset_wave * route_threads if reset_expert < num_experts: route_counts[src_rank[0], reset_expert] = 0 - T.sync_threads(7, 64) + T.sync_threads(7, route_threads) - for assign_wave in T.serial(ceil_div(num_routes, 64)): - assign_route = tx + assign_wave * 64 + for assign_wave in T.serial(ceil_div(num_routes, route_threads)): + assign_route = tx + assign_wave * route_threads if assign_route < num_routes: assign_token = assign_route // num_topk assign_topk = assign_route % num_topk @@ -372,10 +373,10 @@ def main( ) else: route_slots[assign_token, assign_topk] = -1 - T.sync_threads(7, 64) + T.sync_threads(7, route_threads) - for publish_wave in T.serial(ceil_div(num_experts * num_ranks, 64)): - publish_idx = tx + publish_wave * 64 + for publish_wave in T.serial(ceil_div(num_experts * num_ranks, route_threads)): + publish_idx = tx + publish_wave * route_threads if publish_idx < num_experts * num_ranks: publish_rank = publish_idx // num_experts publish_expert = publish_idx % num_experts @@ -388,7 +389,7 @@ def main( T.barrier_blocks(barrier[0]) - if tx < 64: + if tx < route_threads: if tx < num_experts_per_rank: recv_count = T.alloc_var(T.int32, init=0) recv_expert = src_rank[0] * num_experts_per_rank + tx @@ -396,8 +397,8 @@ def main( recv_count += route_counts[count_rank, recv_expert] recv_counts[tx] = recv_count - for prefix_wave in T.serial(ceil_div(num_routes, 64)): - prefix_route = tx + prefix_wave * 64 + for prefix_wave in T.serial(ceil_div(num_routes, route_threads)): + prefix_route = tx + prefix_wave * route_threads if prefix_route < num_routes: prefix_token = prefix_route // num_topk prefix_topk = prefix_route % num_topk From b56d0d96c5c93eaa384f1fcf8ea13fd69f2eebb8 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Wed, 12 Aug 2026 17:52:06 +0800 Subject: [PATCH 11/30] feat(distributed): fuse Pro mega MoE path --- .../mega_moe/example_sm90_fp8_mega_moe.py | 107 ++++++++++++------ 1 file changed, 74 insertions(+), 33 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 97cfd77fc3..13b90a2f62 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -1,9 +1,8 @@ """Multi-GPU FP8 Mega MoE for SM90 using TileScale distributed primitives. -The Flash configuration uses two persistent kernels: one for routing, dispatch, -L1 GEMM, and SwiGLU quantization, and one for L2 GEMM, scatter, and reduction. -The Pro configuration retains the multi-kernel path until its four-warpgroup L1 -kernel is ready. +The fused implementation uses two persistent kernels: one for routing, +dispatch, L1 GEMM, and SwiGLU quantization, and one for L2 GEMM, scatter, and +reduction. Both model configurations use manually selected warp counts. """ from __future__ import annotations @@ -302,6 +301,13 @@ def fused_l1_swiglu_manual_warp_kernel( num_n_blocks = ceil_div(l1_n, block_n) num_compute_tiles = num_experts_per_rank * num_m_blocks * num_n_blocks num_k_blocks = hidden // block_k + num_math_threads = threads - 128 + num_output_scale_groups = block_n // (2 * SCALE_GRANULARITY) + num_l1_scale_groups = l1_n // (2 * SCALE_GRANULARITY) + tma_block_n = min(block_n, 256) + num_tma_n_blocks = block_n // tma_block_n + frontend_registers = 32 if num_math_threads == 512 else 48 + math_registers = 112 if num_math_threads == 512 else 208 dispatch_thread = 0 route_threads = 256 @@ -348,7 +354,9 @@ def main( T.float8_e4m3fn, ) out_shared = T.alloc_shared((block_m, block_n // 2), T.float8_e4m3fn) - stage_barriers = T.alloc_barrier([64] * pipeline_stages + [256] * pipeline_stages) + stage_barriers = T.alloc_barrier( + [64] * pipeline_stages + [num_math_threads] * pipeline_stages + ) if bid == 0: if tx < route_threads: @@ -416,9 +424,9 @@ def main( T.sync_grid() if tx < 128: - T.dec_max_nreg(48) + T.dec_max_nreg(frontend_registers) else: - T.inc_max_nreg(208) + T.inc_max_nreg(math_registers) dispatch_warp = tx // 32 dispatch_lane = tx % 32 @@ -532,15 +540,23 @@ def main( a_shared[producer_stage, :, :], barrier=stage_barriers[producer_stage], ) - T.tma_copy( - l1_weight[ - producer_expert, - producer_n * block_n : (producer_n + 1) * block_n, - producer_k * block_k : (producer_k + 1) * block_k, - ], - b_shared[producer_stage, :, :], - barrier=stage_barriers[producer_stage], - ) + for producer_n_block in T.serial(num_tma_n_blocks): + T.tma_copy( + l1_weight[ + producer_expert, + producer_n * block_n + + producer_n_block * tma_block_n : producer_n * block_n + + (producer_n_block + 1) * tma_block_n, + producer_k * block_k : (producer_k + 1) * block_k, + ], + b_shared[ + producer_stage, + producer_n_block + * tma_block_n : (producer_n_block + 1) * tma_block_n, + :, + ], + barrier=stage_barriers[producer_stage], + ) T.mbarrier_arrive(stage_barriers[producer_stage]) producer_step += num_k_blocks @@ -548,12 +564,16 @@ def main( partial = T.alloc_fragment((block_m, block_n), T.float32) accum = T.alloc_fragment((block_m, block_n), T.bfloat16) gate = T.alloc_fragment((block_m, block_n // 2), T.float32) + gate_grouped = T.reshape( + gate, + (block_m, num_output_scale_groups, SCALE_GRANULARITY), + ) up = T.alloc_fragment((block_m, block_n // 2), T.float32) - amax = T.alloc_fragment((block_m,), T.float32) - scale = T.alloc_fragment((block_m,), T.float32) + amax = T.alloc_fragment((block_m, num_output_scale_groups), T.float32) + scale = T.alloc_fragment((block_m, num_output_scale_groups), T.float32) quant_fp8 = T.alloc_fragment((block_m, block_n // 2), T.float8_e4m3fn) act_scale = T.alloc_fragment((block_m,), T.float32) - weight_scale = T.alloc_local((2,), T.float32) + weight_scale = T.alloc_local((2 * num_output_scale_groups,), T.float32) consumer_step = T.alloc_var(T.int32, init=0) for consumer_wave in T.serial(ceil_div(num_compute_tiles, num_sms)): @@ -584,17 +604,28 @@ def main( consumer_m * block_m + i, consumer_k, ] - weight_scale[0] = l1_weight_sf[consumer_expert, consumer_n, consumer_k] - weight_scale[1] = l1_weight_sf[ - consumer_expert, - consumer_n + num_n_blocks, - consumer_k, - ] + for scale_group in T.serial(num_output_scale_groups): + weight_scale[2 * scale_group] = l1_weight_sf[ + consumer_expert, + consumer_n * num_output_scale_groups + scale_group, + consumer_k, + ] + weight_scale[2 * scale_group + 1] = l1_weight_sf[ + consumer_expert, + num_l1_scale_groups + + consumer_n * num_output_scale_groups + + scale_group, + consumer_k, + ] for i, j in T.Parallel(block_m, block_n): accum[i, j] = ( T.cast(partial[i, j], T.bfloat16) * T.cast( - act_scale[i] * weight_scale[(j % 16) // 8], + act_scale[i] + * weight_scale[ + 2 * (j // (2 * SCALE_GRANULARITY)) + + (j % 16) // 8 + ], T.bfloat16, ) + accum[i, j] @@ -619,17 +650,20 @@ def main( consumer_m * block_m + i, ] ) - T.reduce_absmax(gate, amax, dim=1) - for i in T.Parallel(block_m): - scale[i] = T.max(amax[i], 1e-4) / FP8_MAX + T.reduce_absmax(gate_grouped, amax, dim=2) + for i, scale_group in T.Parallel( + block_m, + num_output_scale_groups, + ): + scale[i, scale_group] = T.max(amax[i, scale_group], 1e-4) / FP8_MAX l2_x_sf[ consumer_expert, consumer_m * block_m + i, - consumer_n, - ] = scale[i] + consumer_n * num_output_scale_groups + scale_group, + ] = scale[i, scale_group] for i, j in T.Parallel(block_m, block_n // 2): gate[i, j] = T.clamp( - gate[i, j] / scale[i], + gate[i, j] / scale[i, j // SCALE_GRANULARITY], -FP8_MAX, FP8_MAX, ) @@ -1160,7 +1194,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): num_topk = model["num_topk"] num_tokens = args.num_tokens activation_clamp = args.activation_clamp - use_fused = args.model_config != "pro" + use_fused = args.implementation == "fused" assert num_experts % num_local_ranks == 0 assert hidden % 256 == 0 and intermediate_hidden % 128 == 0 @@ -1189,6 +1223,11 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): ) if use_fused: + l1_config = ( + {"block_n": 256, "threads": 384, "pipeline_stages": 4} + if args.model_config == "pro" + else {} + ) kernel_specs = [ fused_l1_swiglu_manual_warp_kernel( num_tokens, @@ -1200,6 +1239,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): capacity, num_sms, activation_clamp=activation_clamp, + **l1_config, ), fused_l2_scatter_reduce_manual_warp_kernel( num_tokens, @@ -1470,6 +1510,7 @@ def run_pipeline(check_capacity: bool = False): parser = argparse.ArgumentParser() parser.add_argument("--num-processes", type=int, default=8) parser.add_argument("--model-config", choices=tuple(MODEL_CONFIGS), default="smoke") + parser.add_argument("--implementation", choices=("fused", "multi-kernel"), default="fused") parser.add_argument("--num-tokens", type=int, default=64) parser.add_argument("--capacity", type=int, default=None) parser.add_argument("--activation-clamp", type=float, default=10.0) From 3fa62ba03c4ab78ccd8f68cea759a1929e1b46ea Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Thu, 13 Aug 2026 10:10:50 +0800 Subject: [PATCH 12/30] test(distributed): select fused mega MoE path --- examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py | 1 + 1 file changed, 1 insertion(+) diff --git a/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py index c70f1b3e4a..ddb5e54ae4 100644 --- a/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py @@ -13,6 +13,7 @@ def test_example_sm90_fp8_mega_moe(local_rank: int, num_ranks: int): args = argparse.Namespace( num_processes=num_ranks, model_config="smoke", + implementation="fused", num_tokens=32, capacity=64, activation_clamp=10.0, From 23369cfd1404128eab6eccea400e04ce2f6534e0 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Thu, 13 Aug 2026 10:30:27 +0800 Subject: [PATCH 13/30] refactor(distributed): remove mega MoE fallback kernels --- .../mega_moe/example_sm90_fp8_mega_moe.py | 581 +++--------------- .../test_example_sm90_fp8_mega_moe.py | 1 - 2 files changed, 77 insertions(+), 505 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 13b90a2f62..c6a15fc964 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -91,192 +91,6 @@ def dequantize_block(x: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: return (x_view * scale.unsqueeze(-1).unsqueeze(-3)).view(groups, n, k) -def assign_local_routes_kernel( - num_tokens: int, - num_experts: int, - num_topk: int, - num_ranks: int, - threads: int = 128, -): - @T.prim_func - def main( - topk_idx: T.Tensor((num_tokens, num_topk), T.int32), - route_counts: T.Tensor((num_ranks, num_experts), T.int32), - route_slots: T.Tensor((num_tokens, num_topk), T.int32), - ): - with T.Kernel(T.ceildiv(num_tokens * num_topk, threads), threads=threads) as bx: - src_rank = T.alloc_local((1,), T.int32) - src_rank[0] = T.get_rank() - route_idx = bx * threads + T.get_thread_binding() - if src_rank[0] < num_ranks and route_idx < num_tokens * num_topk: - token_idx = route_idx // num_topk - topk_slot = route_idx % num_topk - expert_idx = topk_idx[token_idx, topk_slot] - if expert_idx >= 0 and expert_idx < num_experts: - route_slots[token_idx, topk_slot] = T.atomic_add( - route_counts[src_rank[0], expert_idx], - 1, - memory_order="relaxed", - return_prev=True, - ) - else: - route_slots[token_idx, topk_slot] = -1 - - return main - - -def reset_route_counts_kernel(num_experts: int, num_ranks: int, threads: int = 128): - @T.prim_func - def main(route_counts: T.Tensor((num_ranks, num_experts), T.int32)): - with T.Kernel(T.ceildiv(num_experts, threads), threads=threads) as bx: - src_rank = T.alloc_local((1,), T.int32) - src_rank[0] = T.get_rank() - expert_idx = bx * threads + T.get_thread_binding() - if src_rank[0] < num_ranks and expert_idx < num_experts: - route_counts[src_rank[0], expert_idx] = 0 - - return main - - -def publish_route_counts_kernel(num_experts: int, num_ranks: int, threads: int = 128): - @T.prim_func - def main(route_counts: T.Tensor((num_ranks, num_experts), T.int32)): - with T.Kernel(T.ceildiv(num_experts * num_ranks, threads), threads=threads) as bx: - src_rank = T.alloc_local((1,), T.int32) - src_rank[0] = T.get_rank() - idx = bx * threads + T.get_thread_binding() - if idx < num_experts * num_ranks: - dst_rank = idx // num_experts - expert_idx = idx % num_experts - if dst_rank != src_rank[0]: - T.st( - route_counts[src_rank[0], expert_idx], - route_counts[src_rank[0], expert_idx], - dst_pe=dst_rank, - ) - - return main - - -def device_barrier_kernel(num_ranks: int): - @T.prim_func - def main(barrier: T.Tensor((num_ranks,), T.int32)): - with T.Kernel(1, threads=32): - rank = T.alloc_local((1,), T.int32) - rank[0] = T.get_rank() - if rank[0] < num_ranks: - T.barrier_blocks(barrier[0]) - T.fence_sys() - - return main - - -def finalize_routes_kernel( - num_tokens: int, - num_experts: int, - num_topk: int, - num_ranks: int, - threads: int = 128, -): - num_experts_per_rank = num_experts // num_ranks - num_routes = num_tokens * num_topk - num_work_items = max(num_routes, num_experts_per_rank) - - @T.prim_func - def main( - topk_idx: T.Tensor((num_tokens, num_topk), T.int32), - route_counts: T.Tensor((num_ranks, num_experts), T.int32), - recv_counts: T.Tensor((num_experts_per_rank,), T.int32), - route_slots: T.Tensor((num_tokens, num_topk), T.int32), - ): - with T.Kernel(T.ceildiv(num_work_items, threads), threads=threads) as bx: - src_rank = T.alloc_local((1,), T.int32) - src_rank[0] = T.get_rank() - work_idx = bx * threads + T.get_thread_binding() - - if work_idx < num_experts_per_rank: - count = T.alloc_var(T.int32, init=0) - expert_idx = src_rank[0] * num_experts_per_rank + work_idx - for peer_rank in T.serial(num_ranks): - count += route_counts[peer_rank, expert_idx] - recv_counts[work_idx] = count - - if work_idx < num_routes: - token_idx = work_idx // num_topk - topk_slot = work_idx % num_topk - expert_idx = topk_idx[token_idx, topk_slot] - slot = T.alloc_var(T.int32, init=route_slots[token_idx, topk_slot]) - if token_idx < num_tokens and expert_idx >= 0 and expert_idx < num_experts and slot >= 0: - for peer_rank in T.serial(num_ranks): - if peer_rank < src_rank[0]: - slot += route_counts[peer_rank, expert_idx] - route_slots[token_idx, topk_slot] = slot - - return main - - -def dispatch_tokens_kernel( - num_tokens: int, - hidden: int, - num_experts: int, - num_topk: int, - num_ranks: int, - capacity: int, - block_h: int = 256, - threads: int = 128, -): - num_experts_per_rank = num_experts // num_ranks - num_scale_groups = hidden // SCALE_GRANULARITY - - @T.prim_func - def main( - x: T.Tensor((num_tokens, hidden), T.float8_e4m3fn), - x_sf: T.Tensor((num_tokens, num_scale_groups), T.float32), - topk_idx: T.Tensor((num_tokens, num_topk), T.int32), - topk_weights: T.Tensor((num_tokens, num_topk), T.float32), - route_slots: T.Tensor((num_tokens, num_topk), T.int32), - recv_x: T.Tensor((num_experts_per_rank, capacity, hidden), T.float8_e4m3fn), - recv_x_sf: T.Tensor( - (num_experts_per_rank, capacity, num_scale_groups), - T.float32, - ), - recv_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), - src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), - src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), - src_topk: T.Tensor((num_experts_per_rank, capacity), T.int32), - ): - with T.Kernel(T.ceildiv(hidden, block_h), num_tokens * num_topk, threads=threads) as (bx, by): - src_rank = T.alloc_local((1,), T.int32) - src_rank[0] = T.get_rank() - tx = T.get_thread_binding() - token_idx = by // num_topk - topk_slot = by % num_topk - expert_idx = topk_idx[token_idx, topk_slot] - slot = route_slots[token_idx, topk_slot] - if expert_idx >= 0 and slot >= 0 and slot < capacity: - dst_rank = expert_idx // num_experts_per_rank - local_expert = expert_idx % num_experts_per_rank - T.copy( - x[token_idx, bx * block_h : (bx + 1) * block_h], - recv_x[local_expert, slot, bx * block_h : (bx + 1) * block_h], - dst_pe=dst_rank, - disable_tma=True, - ) - if bx == 0 and num_scale_groups > 0: - T.copy( - x_sf[token_idx, :], - recv_x_sf[local_expert, slot, :], - dst_pe=dst_rank, - disable_tma=True, - ) - if tx == 0 and src_rank[0] < num_ranks: - T.st(recv_weights[local_expert, slot], topk_weights[token_idx, topk_slot], dst_pe=dst_rank) - T.st(src_ranks[local_expert, slot], src_rank[0], dst_pe=dst_rank) - T.st(src_tokens[local_expert, slot], token_idx, dst_pe=dst_rank) - T.st(src_topk[local_expert, slot], topk_slot, dst_pe=dst_rank) - - return main - def fused_l1_swiglu_manual_warp_kernel( num_tokens: int, @@ -681,56 +495,6 @@ def main( return main -def fp8_grouped_gemm_kernel( - num_experts_per_rank: int, - capacity: int, - n: int, - k: int, - block_m: int = 64, - block_n: int = 128, - block_k: int = 128, - threads: int = 128, - pipeline_stages: int = 4, -): - @T.prim_func - def main( - a: T.Tensor((num_experts_per_rank, capacity, k), T.float8_e4m3fn), - b: T.Tensor((num_experts_per_rank, n, k), T.float8_e4m3fn), - a_sf: T.Tensor((num_experts_per_rank, capacity, k // SCALE_GRANULARITY), T.float32), - b_sf: T.Tensor( - (num_experts_per_rank, n // SCALE_GRANULARITY, k // SCALE_GRANULARITY), - T.float32, - ), - recv_counts: T.Tensor((num_experts_per_rank,), T.int32), - out: T.Tensor((num_experts_per_rank, capacity, n), T.bfloat16), - ): - with T.Kernel(T.ceildiv(n, block_n), T.ceildiv(capacity, block_m), num_experts_per_rank, threads=threads) as ( - bx, - by, - bz, - ): - a_shared = T.alloc_shared((block_m, block_k), T.float8_e4m3fn) - b_shared = T.alloc_shared((block_n, block_k), T.float8_e4m3fn) - out_shared = T.alloc_shared((block_m, block_n), T.bfloat16) - partial = T.alloc_fragment((block_m, block_n), T.float32) - accum = T.alloc_fragment((block_m, block_n), T.float32) - - if by * block_m < recv_counts[bz]: - T.clear(partial) - T.clear(accum) - for ko in T.Pipelined(k // block_k, num_stages=pipeline_stages): - T.copy(a[bz, by * block_m, ko * block_k], a_shared) - T.copy(b[bz, bx * block_n, ko * block_k], b_shared) - T.gemm(a_shared, b_shared, partial, transpose_B=True) - b_scale = b_sf[bz, bx, ko] - for i, j in T.Parallel(block_m, block_n): - accum[i, j] += partial[i, j] * (a_sf[bz, by * block_m + i, ko] * b_scale) - T.clear(partial) - T.copy(accum, out_shared) - T.copy(out_shared, out[bz, by * block_m, bx * block_n]) - - return main - def fused_l2_scatter_reduce_manual_warp_kernel( num_tokens: int, @@ -981,119 +745,6 @@ def main( return main -def swiglu_quant_kernel( - num_experts_per_rank: int, - capacity: int, - intermediate_hidden: int, - block_m: int = 8, - block_n: int = 128, - threads: int = 128, - activation_clamp: float = 10.0, -): - @T.prim_func - def main( - gate_up: T.Tensor((num_experts_per_rank, capacity, 2 * intermediate_hidden), T.bfloat16), - route_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), - recv_counts: T.Tensor((num_experts_per_rank,), T.int32), - out: T.Tensor((num_experts_per_rank, capacity, intermediate_hidden), T.float8_e4m3fn), - out_sf: T.Tensor( - (num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), - T.float32, - ), - ): - with T.Kernel( - intermediate_hidden // block_n, - T.ceildiv(capacity, block_m), - num_experts_per_rank, - threads=threads, - ) as (bx, by, bz): - gate = T.alloc_fragment((block_m, block_n), T.float32) - up = T.alloc_fragment((block_m, block_n), T.float32) - activated = T.alloc_fragment((block_m, block_n), T.float32) - amax = T.alloc_fragment((block_m,), T.float32) - scale = T.alloc_fragment((block_m,), T.float32) - quant = T.alloc_fragment((block_m, block_n), T.float32) - quant_fp8 = T.alloc_fragment((block_m, block_n), T.float8_e4m3fn) - - if by * block_m < recv_counts[bz]: - T.copy(gate_up[bz, by * block_m, bx * block_n], gate) - T.copy(gate_up[bz, by * block_m, intermediate_hidden + bx * block_n], up) - for i, j in T.Parallel(block_m, block_n): - gate[i, j] = T.min(gate[i, j], activation_clamp) - up[i, j] = T.max(T.min(up[i, j], activation_clamp), -activation_clamp) - activated[i, j] = gate[i, j] * T.sigmoid(gate[i, j]) * up[i, j] * route_weights[bz, by * block_m + i] - T.reduce_absmax(activated, amax, dim=1) - for i in T.Parallel(block_m): - scale[i] = T.max(amax[i], 1e-4) / FP8_MAX - out_sf[bz, by * block_m + i, bx] = scale[i] - for i, j in T.Parallel(block_m, block_n): - quant[i, j] = T.clamp(activated[i, j] / scale[i], -FP8_MAX, FP8_MAX) - T.copy(quant, quant_fp8) - T.copy(quant_fp8, out[bz, by * block_m, bx * block_n]) - - return main - - -def scatter_outputs_kernel( - num_experts_per_rank: int, - capacity: int, - num_tokens: int, - num_topk: int, - hidden: int, - block_h: int = 256, - threads: int = 128, -): - @T.prim_func - def main( - local_out: T.Tensor((num_experts_per_rank, capacity, hidden), T.bfloat16), - recv_counts: T.Tensor((num_experts_per_rank,), T.int32), - src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), - src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), - src_topk: T.Tensor((num_experts_per_rank, capacity), T.int32), - combine: T.Tensor((num_tokens, num_topk, hidden), T.bfloat16), - ): - with T.Kernel(T.ceildiv(hidden, block_h), capacity, num_experts_per_rank, threads=threads) as (bx, by, bz): - if by < recv_counts[bz]: - dst_rank = src_ranks[bz, by] - token_idx = src_tokens[bz, by] - topk_slot = src_topk[bz, by] - if dst_rank >= 0 and token_idx < num_tokens and topk_slot < num_topk: - T.copy( - local_out[bz, by, bx * block_h : (bx + 1) * block_h], - combine[token_idx, topk_slot, bx * block_h : (bx + 1) * block_h], - dst_pe=dst_rank, - disable_tma=True, - ) - - return main - - -def reduce_topk_kernel( - num_tokens: int, - num_topk: int, - hidden: int, - block_m: int = 8, - block_h: int = 128, - threads: int = 128, -): - @T.prim_func - def main( - combine: T.Tensor((num_tokens, num_topk, hidden), T.bfloat16), - out: T.Tensor((num_tokens, hidden), T.bfloat16), - ): - with T.Kernel(T.ceildiv(hidden, block_h), T.ceildiv(num_tokens, block_m), threads=threads) as (bx, by): - accum = T.alloc_fragment((block_m, block_h), T.float32) - out_shared = T.alloc_shared((block_m, block_h), T.bfloat16) - T.clear(accum) - for topk_slot in T.serial(num_topk): - for i, j in T.Parallel(block_m, block_h): - if by * block_m + i < num_tokens: - accum[i, j] += combine[by * block_m + i, topk_slot, bx * block_h + j] - T.copy(accum, out_shared) - T.copy(out_shared, out[by * block_m, bx * block_h]) - - return main - def _allocator_size_bytes( num_tokens: int, @@ -1118,13 +769,16 @@ def _allocator_size_bytes( hidden * fp8 + (hidden // 128) * fp32 + 4 * i32 - + 2 * intermediate_hidden * bf16 + intermediate_hidden * fp8 + (intermediate_hidden // 128) * fp32 - + hidden * bf16 ) ) - input_bytes = num_tokens * (hidden * fp8 + (hidden // 128) * fp32 + num_topk * (3 * i32 + fp32) + num_topk * hidden * bf16) + input_bytes = num_tokens * ( + hidden * fp8 + + (hidden // 128) * fp32 + + num_topk * (3 * i32 + fp32) + + (num_topk + 1) * hidden * bf16 + ) return align_up(weight_bytes + weight_scale_bytes + pool_bytes + input_bytes + 2**27, 2**20) @@ -1194,7 +848,6 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): num_topk = model["num_topk"] num_tokens = args.num_tokens activation_clamp = args.activation_clamp - use_fused = args.implementation == "fused" assert num_experts % num_local_ranks == 0 assert hidden % 256 == 0 and intermediate_hidden % 128 == 0 @@ -1222,74 +875,39 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): use_vmm=True, ) - if use_fused: - l1_config = ( - {"block_n": 256, "threads": 384, "pipeline_stages": 4} - if args.model_config == "pro" - else {} - ) - kernel_specs = [ - fused_l1_swiglu_manual_warp_kernel( - num_tokens, - hidden, - 2 * intermediate_hidden, - num_experts, - num_topk, - num_ranks, - capacity, - num_sms, - activation_clamp=activation_clamp, - **l1_config, - ), - fused_l2_scatter_reduce_manual_warp_kernel( - num_tokens, - hidden, - intermediate_hidden, - num_experts_per_rank, - num_topk, - num_ranks, - capacity, - num_sms, - ), - ] - else: - kernel_specs = [ - reset_route_counts_kernel(num_experts, num_ranks), - assign_local_routes_kernel(num_tokens, num_experts, num_topk, num_ranks), - publish_route_counts_kernel(num_experts, num_ranks), - device_barrier_kernel(num_ranks), - finalize_routes_kernel(num_tokens, num_experts, num_topk, num_ranks), - dispatch_tokens_kernel(num_tokens, hidden, num_experts, num_topk, num_ranks, capacity), - fp8_grouped_gemm_kernel(num_experts_per_rank, capacity, 2 * intermediate_hidden, hidden), - swiglu_quant_kernel( - num_experts_per_rank, - capacity, - intermediate_hidden, - activation_clamp=activation_clamp, - ), - fp8_grouped_gemm_kernel(num_experts_per_rank, capacity, hidden, intermediate_hidden), - scatter_outputs_kernel(num_experts_per_rank, capacity, num_tokens, num_topk, hidden), - reduce_topk_kernel(num_tokens, num_topk, hidden), - ] + l1_config = ( + {"block_n": 256, "threads": 384, "pipeline_stages": 4} + if args.model_config == "pro" + else {} + ) + kernel_specs = [ + fused_l1_swiglu_manual_warp_kernel( + num_tokens, + hidden, + 2 * intermediate_hidden, + num_experts, + num_topk, + num_ranks, + capacity, + num_sms, + activation_clamp=activation_clamp, + **l1_config, + ), + fused_l2_scatter_reduce_manual_warp_kernel( + num_tokens, + hidden, + intermediate_hidden, + num_experts_per_rank, + num_topk, + num_ranks, + capacity, + num_sms, + ), + ] kernels = [tilelang.compile(spec, compile_once=True, compile_group=group) for spec in kernel_specs] for kernel in kernels: kernel.initialize(allocator=allocator) - if use_fused: - fused_l1, fused_l2 = kernels - else: - ( - reset_route_counts, - assign_local_routes, - publish_route_counts, - device_barrier, - finalize_routes, - dispatch_tokens, - l1_gemm, - swiglu_quant, - l2_gemm, - scatter_outputs, - reduce_topk, - ) = kernels + fused_l1, fused_l2 = kernels if local_rank == 0 and args.print_source: for kernel in kernels: @@ -1329,7 +947,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): x_sf = allocator_tensor(x_sf_src.shape, x_sf_src.dtype, allocator=allocator).copy_(x_sf_src) topk_idx = allocator_tensor(topk_idx_src.shape, topk_idx_src.dtype, allocator=allocator).copy_(topk_idx_src) topk_weights = allocator_tensor(topk_weights_src.shape, topk_weights_src.dtype, allocator=allocator).copy_(topk_weights_src) - l1_fp8_kernel = interleave_gate_up_weights(l1_fp8_src) if use_fused else l1_fp8_src + l1_fp8_kernel = interleave_gate_up_weights(l1_fp8_src) l1_fp8 = allocator_tensor(l1_fp8_kernel.shape, l1_fp8_kernel.dtype, allocator=allocator).copy_(l1_fp8_kernel) del l1_fp8_kernel l1_sf = allocator_tensor(l1_sf_src.shape, l1_sf_src.dtype, allocator=allocator).copy_(l1_sf_src) @@ -1349,30 +967,17 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): src_tokens = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) src_topk = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) route_slots = allocator_tensor((num_tokens, num_topk), torch.int32, allocator=allocator) - if use_fused: - arrivals = allocator_tensor( - (num_experts_per_rank, ceil_div(capacity, 64)), - torch.uint32, - allocator=allocator, - ) - if not use_fused: - l1_out = allocator_tensor( - (num_experts_per_rank, capacity, 2 * intermediate_hidden), - torch.bfloat16, - allocator=allocator, - ) + arrivals = allocator_tensor( + (num_experts_per_rank, ceil_div(capacity, 64)), + torch.uint32, + allocator=allocator, + ) l2_x = allocator_tensor((num_experts_per_rank, capacity, intermediate_hidden), torch.float8_e4m3fn, allocator=allocator) l2_x_sf = allocator_tensor( (num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), torch.float32, allocator=allocator, ) - if not use_fused: - l2_out = allocator_tensor( - (num_experts_per_rank, capacity, hidden), - torch.bfloat16, - allocator=allocator, - ) combine = allocator_tensor((num_tokens, num_topk, hidden), torch.bfloat16, allocator=allocator) out = allocator_tensor((num_tokens, hidden), torch.bfloat16, allocator=allocator) @@ -1380,8 +985,7 @@ def reset_state(): route_counts.zero_() barrier.zero_() recv_counts.zero_() - if use_fused: - arrivals.zero_() + arrivals.zero_() recv_x.zero_() recv_x_sf.zero_() recv_weights.zero_() @@ -1391,73 +995,44 @@ def reset_state(): dist.barrier(group=group) def run_pipeline(check_capacity: bool = False): - if use_fused: - fused_l1( - x, - x_sf, - topk_idx, - topk_weights, - route_counts, - recv_counts, - route_slots, - arrivals, - recv_x, - recv_x_sf, - recv_weights, - src_ranks, - src_tokens, - src_topk, - l1_fp8, - l1_sf, - l2_x, - l2_x_sf, - barrier, - ) - else: - reset_route_counts(route_counts) - assign_local_routes(topk_idx, route_counts, route_slots) - publish_route_counts(route_counts) - device_barrier(barrier) - finalize_routes(topk_idx, route_counts, recv_counts, route_slots) - dispatch_tokens( - x, - x_sf, - topk_idx, - topk_weights, - route_slots, - recv_x, - recv_x_sf, - recv_weights, - src_ranks, - src_tokens, - src_topk, - ) - device_barrier(barrier) + fused_l1( + x, + x_sf, + topk_idx, + topk_weights, + route_counts, + recv_counts, + route_slots, + arrivals, + recv_x, + recv_x_sf, + recv_weights, + src_ranks, + src_tokens, + src_topk, + l1_fp8, + l1_sf, + l2_x, + l2_x_sf, + barrier, + ) if check_capacity: local_max = recv_counts.max() dist.all_reduce(local_max, op=dist.ReduceOp.MAX, group=group) assert local_max.item() <= capacity, f"expert capacity {capacity} is smaller than received routes {local_max.item()}" - if use_fused: - fused_l2( - l2_x, - l2_fp8, - l2_x_sf, - l2_sf, - recv_counts, - src_ranks, - src_tokens, - src_topk, - combine, - barrier, - out, - ) - else: - l1_gemm(recv_x, l1_fp8, recv_x_sf, l1_sf, recv_counts, l1_out) - swiglu_quant(l1_out, recv_weights, recv_counts, l2_x, l2_x_sf) - l2_gemm(l2_x, l2_fp8, l2_x_sf, l2_sf, recv_counts, l2_out) - scatter_outputs(l2_out, recv_counts, src_ranks, src_tokens, src_topk, combine) - device_barrier(barrier) - reduce_topk(combine, out) + fused_l2( + l2_x, + l2_fp8, + l2_x_sf, + l2_sf, + recv_counts, + src_ranks, + src_tokens, + src_topk, + combine, + barrier, + out, + ) return out reset_state() @@ -1498,7 +1073,6 @@ def run_pipeline(check_capacity: bool = False): if local_rank == 0: print( f"tilescale sm90 fp8 mega moe: model={args.model_config} M={num_tokens} " - f"implementation={'fused' if use_fused else 'multi-kernel'} " f"capacity={capacity} latency={latency * 1000:.1f} us" ) @@ -1510,7 +1084,6 @@ def run_pipeline(check_capacity: bool = False): parser = argparse.ArgumentParser() parser.add_argument("--num-processes", type=int, default=8) parser.add_argument("--model-config", choices=tuple(MODEL_CONFIGS), default="smoke") - parser.add_argument("--implementation", choices=("fused", "multi-kernel"), default="fused") parser.add_argument("--num-tokens", type=int, default=64) parser.add_argument("--capacity", type=int, default=None) parser.add_argument("--activation-clamp", type=float, default=10.0) diff --git a/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py index ddb5e54ae4..c70f1b3e4a 100644 --- a/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py @@ -13,7 +13,6 @@ def test_example_sm90_fp8_mega_moe(local_rank: int, num_ranks: int): args = argparse.Namespace( num_processes=num_ranks, model_config="smoke", - implementation="fused", num_tokens=32, capacity=64, activation_clamp=10.0, From 1c6c03d6cac65472aa13a36ee9cd5888e6de9c15 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Thu, 13 Aug 2026 11:10:05 +0800 Subject: [PATCH 14/30] perf(distributed): broadcast mega MoE scatter metadata --- .../mega_moe/example_sm90_fp8_mega_moe.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index c6a15fc964..04357972de 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -623,6 +623,9 @@ def main( act_scale = T.alloc_fragment((block_m,), T.float32) weight_scale = T.alloc_local((2,), T.float32) consumer_step = T.alloc_var(T.int32, init=0) + scatter_dst_rank = T.alloc_var(T.int32, init=0) + scatter_dst_token = T.alloc_var(T.int32, init=0) + scatter_dst_topk = T.alloc_var(T.int32, init=0) for consumer_wave in T.serial(ceil_div(num_compute_tiles, num_sms)): consumer_tile = bid + consumer_wave * num_sms @@ -686,9 +689,13 @@ def main( row = scatter_warp * (block_m // 8) + row_in_warp pool_row = consumer_m * block_m + row if pool_row < recv_counts[consumer_expert]: - dst_rank = src_ranks[consumer_expert, pool_row] - dst_token = src_tokens[consumer_expert, pool_row] - dst_topk = src_topk[consumer_expert, pool_row] + if tx % 32 == 0: + scatter_dst_rank = src_ranks[consumer_expert, pool_row] + scatter_dst_token = src_tokens[consumer_expert, pool_row] + scatter_dst_topk = src_topk[consumer_expert, pool_row] + dst_rank = T.shfl_sync(scatter_dst_rank, 0) + dst_token = T.shfl_sync(scatter_dst_token, 0) + dst_topk = T.shfl_sync(scatter_dst_topk, 0) if ( dst_rank >= 0 and dst_rank < num_ranks From 888c36f0fc07f0c6e63bbfa699438f9e785f5bc4 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Thu, 13 Aug 2026 16:26:57 +0800 Subject: [PATCH 15/30] feat(distributed): generalize SM90 mega MoE shapes --- .../mega_moe/example_sm90_fp8_mega_moe.py | 125 +++++++++++++++--- .../test_example_sm90_fp8_mega_moe.py | 49 +++++++ 2 files changed, 152 insertions(+), 22 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 04357972de..631e29a9a9 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -2,7 +2,9 @@ The fused implementation uses two persistent kernels: one for routing, dispatch, L1 GEMM, and SwiGLU quantization, and one for L2 GEMM, scatter, and -reduction. Both model configurations use manually selected warp counts. +reduction. Model presets and aligned custom shapes use manually selected warp +counts and shape/load-based pipeline depths. Custom hidden sizes must be at +least 512 and divisible by 256; intermediate sizes must be divisible by 128. """ from __future__ import annotations @@ -42,6 +44,64 @@ def align_up(x: int, alignment: int) -> int: return ceil_div(x, alignment) * alignment +def resolve_model_config(args: argparse.Namespace) -> Tuple[str, dict[str, int]]: + model = MODEL_CONFIGS[args.model_config].copy() + overrides = { + "hidden": getattr(args, "hidden", None), + "intermediate_hidden": getattr(args, "intermediate_hidden", None), + "num_experts": getattr(args, "num_experts", None), + "num_topk": getattr(args, "num_topk", None), + } + is_custom = any(value is not None for value in overrides.values()) + model.update({key: value for key, value in overrides.items() if value is not None}) + return ("custom" if is_custom else args.model_config), model + + +def classify_shape(hidden: int, intermediate_hidden: int) -> str: + if 3072 <= hidden < 5120 and 1536 <= intermediate_hidden < 2560: + return "compact" + if 5120 <= hidden <= 8192 and 2560 <= intermediate_hidden <= 4096: + return "wide" + return "generic" + + +def select_manual_warp_configs( + hidden: int, + intermediate_hidden: int, + num_tokens: int, + num_topk: int, + num_experts_per_rank: int, + num_sms: int, +) -> Tuple[str, dict[str, int], dict[str, int]]: + """Select the TileScale counterpart of DeepGEMM SM90 schedule families.""" + shape_family = classify_shape(hidden, intermediate_hidden) + routed_tokens = num_tokens * num_topk + high_sm = num_sms >= 100 + + l1_stages = 5 + l2_stages = 3 + if high_sm and shape_family == "compact": + if 12 * num_experts_per_rank < routed_tokens <= 32 * num_experts_per_rank: + l1_stages = l2_stages = 3 + elif ( + 128 * num_experts_per_rank < routed_tokens <= 256 * num_experts_per_rank + or routed_tokens > 1024 * num_experts_per_rank + ): + l1_stages = l2_stages = 4 + elif high_sm and shape_family == "wide": + # BN512/BK256 are profitable in the CUDA kernel, but the manually + # tuned TileScale BN256/BK128 path is faster for the current WGMMA + # lowering and remains the generic Wide schedule. + l1_stages = 4 + + common = {"block_m": 64, "block_n": 256, "block_k": 128, "threads": 384} + return ( + shape_family, + {**common, "pipeline_stages": l1_stages}, + {**common, "pipeline_stages": l2_stages}, + ) + + def per_token_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: m, k = x.shape x_view = x.float().view(m, k // SCALE_GRANULARITY, SCALE_GRANULARITY) @@ -212,12 +272,14 @@ def main( T.barrier_blocks(barrier[0]) if tx < route_threads: - if tx < num_experts_per_rank: - recv_count = T.alloc_var(T.int32, init=0) - recv_expert = src_rank[0] * num_experts_per_rank + tx - for count_rank in T.serial(num_ranks): - recv_count += route_counts[count_rank, recv_expert] - recv_counts[tx] = recv_count + for count_wave in T.serial(ceil_div(num_experts_per_rank, route_threads)): + count_local_expert = tx + count_wave * route_threads + if count_local_expert < num_experts_per_rank: + recv_count = T.alloc_var(T.int32, init=0) + recv_expert = src_rank[0] * num_experts_per_rank + count_local_expert + for count_rank in T.serial(num_ranks): + recv_count += route_counts[count_rank, recv_expert] + recv_counts[count_local_expert] = recv_count for prefix_wave in T.serial(ceil_div(num_routes, route_threads)): prefix_route = tx + prefix_wave * route_threads @@ -565,6 +627,7 @@ def main( (pipeline_stages, block_n, block_k), T.float8_e4m3fn, ) + a_sf_shared = T.alloc_shared((pipeline_stages, block_m), T.float32) out_shared = T.alloc_shared((block_m, block_n), T.bfloat16) reduce_shared = T.alloc_shared((reduce_block_m, reduce_block_h), T.bfloat16) stage_barriers = T.alloc_barrier([64] * pipeline_stages + [256] * pipeline_stages) @@ -614,6 +677,11 @@ def main( b_shared[producer_stage, :, :], barrier=stage_barriers[producer_stage], ) + a_sf_shared[producer_stage, tx - 64] = a_sf[ + producer_expert, + producer_m * block_m + tx - 64, + producer_k, + ] T.mbarrier_arrive(stage_barriers[producer_stage]) producer_step += num_k_blocks @@ -656,11 +724,7 @@ def main( transpose_B=True, ) for i in T.Parallel(block_m): - act_scale[i] = a_sf[ - consumer_expert, - consumer_m * block_m + i, - consumer_k, - ] + act_scale[i] = a_sf_shared[consumer_stage, i] weight_scale[0] = b_sf[ consumer_expert, consumer_n * 2, @@ -848,7 +912,7 @@ def allocator_tensor(shape, dtype, allocator): def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): - model = MODEL_CONFIGS[args.model_config] + model_name, model = resolve_model_config(args) hidden = model["hidden"] intermediate_hidden = model["intermediate_hidden"] num_experts = model["num_experts"] @@ -856,11 +920,19 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): num_tokens = args.num_tokens activation_clamp = args.activation_clamp - assert num_experts % num_local_ranks == 0 - assert hidden % 256 == 0 and intermediate_hidden % 128 == 0 + assert num_tokens > 0 + assert hidden >= 512 and hidden % 256 == 0 + assert intermediate_hidden > 0 and intermediate_hidden % 128 == 0 + assert num_experts > 0 and num_experts % num_local_ranks == 0 + assert 0 < num_topk <= min(32, num_experts) num_experts_per_rank = num_experts // num_local_ranks average_recv = ceil_div(num_tokens * num_local_ranks * num_topk, num_experts) - capacity = args.capacity or align_up(max(average_recv * 2, 64), 64) + capacity = ( + args.capacity + if args.capacity is not None + else align_up(max(average_recv * 2, 64), 64) + ) + assert capacity >= 64 and capacity % 64 == 0 rank, num_ranks, group = init_dist(local_rank, num_local_ranks) assert rank == local_rank and num_ranks == num_local_ranks @@ -882,10 +954,13 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): use_vmm=True, ) - l1_config = ( - {"block_n": 256, "threads": 384, "pipeline_stages": 4} - if args.model_config == "pro" - else {} + shape_family, l1_config, l2_config = select_manual_warp_configs( + hidden, + intermediate_hidden, + num_tokens, + num_topk, + num_experts_per_rank, + num_sms, ) kernel_specs = [ fused_l1_swiglu_manual_warp_kernel( @@ -909,6 +984,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): num_ranks, capacity, num_sms, + **l2_config, ), ] kernels = [tilelang.compile(spec, compile_once=True, compile_group=group) for spec in kernel_specs] @@ -1079,8 +1155,9 @@ def run_pipeline(check_capacity: bool = False): ) if local_rank == 0: print( - f"tilescale sm90 fp8 mega moe: model={args.model_config} M={num_tokens} " - f"capacity={capacity} latency={latency * 1000:.1f} us" + f"tilescale sm90 fp8 mega moe: model={model_name} family={shape_family} " + f"M={num_tokens} H={hidden} IH={intermediate_hidden} E={num_experts} " + f"topk={num_topk} capacity={capacity} latency={latency * 1000:.1f} us" ) allocator.close() @@ -1091,6 +1168,10 @@ def run_pipeline(check_capacity: bool = False): parser = argparse.ArgumentParser() parser.add_argument("--num-processes", type=int, default=8) parser.add_argument("--model-config", choices=tuple(MODEL_CONFIGS), default="smoke") + parser.add_argument("--hidden", type=int, default=None) + parser.add_argument("--intermediate-hidden", type=int, default=None) + parser.add_argument("--num-experts", type=int, default=None) + parser.add_argument("--num-topk", type=int, default=None) parser.add_argument("--num-tokens", type=int, default=64) parser.add_argument("--capacity", type=int, default=None) parser.add_argument("--activation-clamp", type=float, default=10.0) diff --git a/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py index c70f1b3e4a..d73cb7618d 100644 --- a/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py @@ -8,6 +8,55 @@ import example_sm90_fp8_mega_moe +def test_custom_model_config_and_schedule(): + args = argparse.Namespace( + model_config="smoke", + hidden=2560, + intermediate_hidden=1536, + num_experts=64, + num_topk=4, + ) + model_name, model = example_sm90_fp8_mega_moe.resolve_model_config(args) + assert model_name == "custom" + assert model == { + "hidden": 2560, + "intermediate_hidden": 1536, + "num_experts": 64, + "num_topk": 4, + } + + family, l1, l2 = example_sm90_fp8_mega_moe.select_manual_warp_configs( + model["hidden"], + model["intermediate_hidden"], + num_tokens=64, + num_topk=model["num_topk"], + num_experts_per_rank=16, + num_sms=132, + ) + assert family == "generic" + assert l1 == { + "block_m": 64, + "block_n": 256, + "block_k": 128, + "threads": 384, + "pipeline_stages": 5, + } + assert l2 == {**l1, "pipeline_stages": 3} + + family, l1, l2 = example_sm90_fp8_mega_moe.select_manual_warp_configs( + 4096, 2048, num_tokens=128, num_topk=6, num_experts_per_rank=32, num_sms=132 + ) + assert family == "compact" + assert l1["pipeline_stages"] == l2["pipeline_stages"] == 3 + + family, l1, l2 = example_sm90_fp8_mega_moe.select_manual_warp_configs( + 7168, 3072, num_tokens=128, num_topk=6, num_experts_per_rank=48, num_sms=132 + ) + assert family == "wide" + assert l1["pipeline_stages"] == 4 + assert l2["pipeline_stages"] == 3 + + @distributed_test(nprocs=4, require_fabric=True) def test_example_sm90_fp8_mega_moe(local_rank: int, num_ranks: int): args = argparse.Namespace( From ef751f576c640784ca1a4812dc537df9b2554813 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Thu, 13 Aug 2026 17:07:43 +0800 Subject: [PATCH 16/30] refactor(distributed): clarify mega MoE warp roles --- .../mega_moe/example_sm90_fp8_mega_moe.py | 102 +++++++++++++----- 1 file changed, 73 insertions(+), 29 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 631e29a9a9..75ecdcc0e5 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -175,14 +175,29 @@ def fused_l1_swiglu_manual_warp_kernel( num_n_blocks = ceil_div(l1_n, block_n) num_compute_tiles = num_experts_per_rank * num_m_blocks * num_n_blocks num_k_blocks = hidden // block_k - num_math_threads = threads - 128 + # CTA roles: warps 0-1 dispatch, warps 2-3 TMA, then WGMMA warpgroups. + warp_size = 32 + warpgroup_size = 128 + dispatch_warps = 2 + producer_warps = 2 + frontend_warps = dispatch_warps + producer_warps + dispatch_threads = dispatch_warps * warp_size + producer_begin = dispatch_threads + producer_threads = producer_warps * warp_size + producer_end = frontend_warps * warp_size + math_begin = producer_end + num_math_threads = threads - math_begin + assert producer_end == producer_begin + producer_threads == warpgroup_size + assert num_math_threads > 0 and num_math_threads % warpgroup_size == 0 + math_warpgroups = num_math_threads // warpgroup_size + assert math_warpgroups in (2, 4) num_output_scale_groups = block_n // (2 * SCALE_GRANULARITY) num_l1_scale_groups = l1_n // (2 * SCALE_GRANULARITY) tma_block_n = min(block_n, 256) num_tma_n_blocks = block_n // tma_block_n frontend_registers = 32 if num_math_threads == 512 else 48 math_registers = 112 if num_math_threads == 512 else 208 - dispatch_thread = 0 + dispatch_leader_lane = 0 route_threads = 256 @T.prim_func @@ -229,7 +244,7 @@ def main( ) out_shared = T.alloc_shared((block_m, block_n // 2), T.float8_e4m3fn) stage_barriers = T.alloc_barrier( - [64] * pipeline_stages + [num_math_threads] * pipeline_stages + [producer_threads] * pipeline_stages + [num_math_threads] * pipeline_stages ) if bid == 0: @@ -299,16 +314,20 @@ def main( T.sync_grid() - if tx < 128: + if tx < math_begin: T.dec_max_nreg(frontend_registers) else: T.inc_max_nreg(math_registers) - dispatch_warp = tx // 32 - dispatch_lane = tx % 32 - if tx < 64: - for metadata_wave in T.serial(ceil_div(num_routes, num_sms * 64)): - metadata_route = bid * 64 + tx + metadata_wave * num_sms * 64 + dispatch_warp = tx // warp_size + dispatch_lane = tx % warp_size + if tx < dispatch_threads: + for metadata_wave in T.serial(ceil_div(num_routes, num_sms * dispatch_threads)): + metadata_route = ( + bid * dispatch_threads + + tx + + metadata_wave * num_sms * dispatch_threads + ) if metadata_route < num_routes: metadata_token = metadata_route // num_topk metadata_topk = metadata_route % num_topk @@ -340,13 +359,16 @@ def main( dst_pe=metadata_rank, ) - if tx < 64: - for pull_wave in T.serial(ceil_div(num_experts_per_rank * capacity, num_sms * 2)): - pull_idx = bid * 2 + dispatch_warp + pull_wave * num_sms * 2 + for pull_wave in T.serial(ceil_div(num_experts_per_rank * capacity, num_sms * dispatch_warps)): + pull_idx = ( + bid * dispatch_warps + + dispatch_warp + + pull_wave * num_sms * dispatch_warps + ) pull_expert = pull_idx // capacity pull_slot = pull_idx % capacity if pull_expert < num_experts_per_rank and pull_slot < recv_counts[pull_expert]: - if dispatch_lane == dispatch_thread: + if dispatch_lane == dispatch_leader_lane: T.wait_ge( src_ranks[pull_expert, pull_slot], 0, @@ -371,7 +393,7 @@ def main( unroll_factor=8, ) T.sync_warp() - if dispatch_lane == dispatch_thread: + if dispatch_lane == dispatch_leader_lane: T.atom_add( arrivals[pull_expert, pull_slot // block_m], 1, @@ -379,7 +401,7 @@ def main( sem="release", ) - if tx >= 64 and tx < 128: + if tx >= producer_begin and tx < producer_end: producer_step = T.alloc_var(T.int32, init=0) for producer_wave in T.serial(ceil_div(num_compute_tiles, num_sms)): producer_tile = bid + producer_wave * num_sms @@ -392,14 +414,14 @@ def main( block_m, recv_counts[producer_expert] - producer_m * block_m, ) - if tx == 64: + if tx == producer_begin: T.wait_ge( arrivals[producer_expert, producer_m], producer_arrivals, scope=T.WaitScope.GPU, semantics=T.WaitSemantics.ACQUIRE, ) - T.sync_threads(5, 64) + T.sync_threads(5, producer_threads) for producer_k in T.serial(num_k_blocks): producer_stage = (producer_step + producer_k) % pipeline_stages producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 @@ -436,7 +458,7 @@ def main( T.mbarrier_arrive(stage_barriers[producer_stage]) producer_step += num_k_blocks - elif tx >= 128: + elif tx >= math_begin: partial = T.alloc_fragment((block_m, block_n), T.float32) accum = T.alloc_fragment((block_m, block_n), T.bfloat16) gate = T.alloc_fragment((block_m, block_n // 2), T.float32) @@ -582,6 +604,26 @@ def fused_l2_scatter_reduce_manual_warp_kernel( num_reduce_n_blocks = ceil_div(hidden, reduce_block_h) num_reduce_m_blocks = ceil_div(num_tokens, reduce_block_m) num_reduce_tiles = num_reduce_n_blocks * num_reduce_m_blocks + # GEMM roles: warps 0-1 idle, warps 2-3 TMA, then two WGMMA + # warpgroups. After the grid barrier, warps 0-3 perform top-k reduction. + warp_size = 32 + warpgroup_size = 128 + reduce_warps = 4 + producer_begin_warp = 2 + producer_warps = 2 + math_begin_warp = 4 + math_warpgroups = 2 + reduce_threads = reduce_warps * warp_size + producer_begin = producer_begin_warp * warp_size + producer_threads = producer_warps * warp_size + producer_end = producer_begin + producer_threads + math_begin = math_begin_warp * warp_size + num_math_threads = math_warpgroups * warpgroup_size + math_warps = num_math_threads // warp_size + rows_per_math_warp = block_m // math_warps + assert producer_end == math_begin == reduce_threads + assert block_m % math_warps == 0 + assert threads == math_begin + num_math_threads @T.prim_func def main( @@ -630,14 +672,16 @@ def main( a_sf_shared = T.alloc_shared((pipeline_stages, block_m), T.float32) out_shared = T.alloc_shared((block_m, block_n), T.bfloat16) reduce_shared = T.alloc_shared((reduce_block_m, reduce_block_h), T.bfloat16) - stage_barriers = T.alloc_barrier([64] * pipeline_stages + [256] * pipeline_stages) + stage_barriers = T.alloc_barrier( + [producer_threads] * pipeline_stages + [num_math_threads] * pipeline_stages + ) - if tx < 128: + if tx < math_begin: T.dec_max_nreg(48) else: T.inc_max_nreg(208) - if tx >= 64 and tx < 128: + if tx >= producer_begin and tx < producer_end: producer_step = T.alloc_var(T.int32, init=0) for producer_wave in T.serial(ceil_div(num_compute_tiles, num_sms)): producer_tile = bid + producer_wave * num_sms @@ -677,15 +721,15 @@ def main( b_shared[producer_stage, :, :], barrier=stage_barriers[producer_stage], ) - a_sf_shared[producer_stage, tx - 64] = a_sf[ + a_sf_shared[producer_stage, tx - producer_begin] = a_sf[ producer_expert, - producer_m * block_m + tx - 64, + producer_m * block_m + tx - producer_begin, producer_k, ] T.mbarrier_arrive(stage_barriers[producer_stage]) producer_step += num_k_blocks - elif tx >= 128: + elif tx >= math_begin: partial = T.alloc_fragment((block_m, block_n), T.float32) accum = T.alloc_fragment((block_m, block_n), T.bfloat16) act_scale = T.alloc_fragment((block_m,), T.float32) @@ -748,12 +792,12 @@ def main( T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) consumer_step += num_k_blocks T.copy(accum, out_shared) - scatter_warp = (tx - 128) // 32 - for row_in_warp in T.serial(block_m // 8): - row = scatter_warp * (block_m // 8) + row_in_warp + scatter_warp = (tx - math_begin) // warp_size + for row_in_warp in T.serial(rows_per_math_warp): + row = scatter_warp * rows_per_math_warp + row_in_warp pool_row = consumer_m * block_m + row if pool_row < recv_counts[consumer_expert]: - if tx % 32 == 0: + if tx % warp_size == 0: scatter_dst_rank = src_ranks[consumer_expert, pool_row] scatter_dst_token = src_tokens[consumer_expert, pool_row] scatter_dst_topk = src_topk[consumer_expert, pool_row] @@ -788,7 +832,7 @@ def main( T.barrier_blocks(barrier[0]) T.sync_grid() - if tx < 128: + if tx < reduce_threads: reduce_accum = T.alloc_fragment((reduce_block_m, reduce_block_h), T.float32) for reduce_wave in T.serial(ceil_div(num_reduce_tiles, num_sms)): reduce_tile = bid + reduce_wave * num_sms From c3805f4beef9ec3872c7d27832e59d10ee435aed Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Thu, 13 Aug 2026 18:29:10 +0800 Subject: [PATCH 17/30] feat(distributed): schedule mega MoE expert waves --- .../mega_moe/example_sm90_fp8_mega_moe.py | 367 ++++++++++++++++-- .../test_example_sm90_fp8_mega_moe.py | 4 + 2 files changed, 332 insertions(+), 39 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 75ecdcc0e5..65aa7695b7 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -65,6 +65,46 @@ def classify_shape(hidden: int, intermediate_hidden: int) -> str: return "generic" +def normalize_experts_per_wave(num_experts: int, requested: int) -> int: + requested = min(max(requested, 1), num_experts) + for candidate in range(requested, num_experts + 1): + if num_experts % candidate == 0: + return candidate + return num_experts + + +def select_generic_experts_per_wave( + intermediate_hidden: int, + routed_tokens: int, + num_experts_per_rank: int, + num_sms: int, + block_m: int = 64, + block_n: int = 256, +) -> int: + if routed_tokens < num_experts_per_rank or routed_tokens > 4 * num_experts_per_rank: + return num_experts_per_rank + + expected_tokens = ceil_div(routed_tokens, num_experts_per_rank) + num_m_blocks = ceil_div(expected_tokens, block_m) + num_n_blocks = 2 * intermediate_hidden // block_n + blocks_per_expert = num_m_blocks * num_n_blocks + requested = min( + num_experts_per_rank, + ceil_div(2 * num_sms, blocks_per_expert), + ) + if blocks_per_expert < num_sms: + max_candidate = min(num_experts_per_rank, 2 * requested) + requested = max( + range(requested, max_candidate + 1), + key=lambda candidate: ( + 1.0 + if num_experts_per_rank % candidate == 0 + else (num_experts_per_rank % candidate) / candidate + ), + ) + return normalize_experts_per_wave(num_experts_per_rank, requested) + + def select_manual_warp_configs( hidden: int, intermediate_hidden: int, @@ -80,25 +120,59 @@ def select_manual_warp_configs( l1_stages = 5 l2_stages = 3 + generic_experts_per_wave = select_generic_experts_per_wave( + intermediate_hidden, + routed_tokens, + num_experts_per_rank, + num_sms, + ) + l1_experts_per_wave = l2_experts_per_wave = generic_experts_per_wave if high_sm and shape_family == "compact": - if 12 * num_experts_per_rank < routed_tokens <= 32 * num_experts_per_rank: + if routed_tokens <= 32 * num_experts_per_rank: l1_stages = l2_stages = 3 + l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave( + num_experts_per_rank, 4 + ) elif ( 128 * num_experts_per_rank < routed_tokens <= 256 * num_experts_per_rank or routed_tokens > 1024 * num_experts_per_rank ): l1_stages = l2_stages = 4 + l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave( + num_experts_per_rank, 32 + ) elif high_sm and shape_family == "wide": # BN512/BK256 are profitable in the CUDA kernel, but the manually # tuned TileScale BN256/BK128 path is faster for the current WGMMA # lowering and remains the generic Wide schedule. l1_stages = 4 + if routed_tokens <= 24 * num_experts_per_rank: + # CUDA selects 16 experts here, while TileScale's direct TIR + # scheduler is faster with a shorter four-expert scan on H200. + l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave( + num_experts_per_rank, 4 + ) + elif 24 * num_experts_per_rank < routed_tokens <= 48 * num_experts_per_rank: + l1_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 8) + l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 48) + elif routed_tokens > 48 * num_experts_per_rank: + l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave( + num_experts_per_rank, 16 + ) common = {"block_m": 64, "block_n": 256, "block_k": 128, "threads": 384} return ( shape_family, - {**common, "pipeline_stages": l1_stages}, - {**common, "pipeline_stages": l2_stages}, + { + **common, + "pipeline_stages": l1_stages, + "num_experts_per_wave": l1_experts_per_wave, + }, + { + **common, + "pipeline_stages": l2_stages, + "num_experts_per_wave": l2_experts_per_wave, + }, ) @@ -167,13 +241,18 @@ def fused_l1_swiglu_manual_warp_kernel( block_k: int = 128, threads: int = 384, pipeline_stages: int = 5, + num_experts_per_wave: int | None = None, ): num_experts_per_rank = num_experts // num_ranks + num_experts_per_wave = num_experts_per_wave or num_experts_per_rank + assert num_experts_per_rank % num_experts_per_wave == 0 num_scale_groups = hidden // SCALE_GRANULARITY num_routes = num_tokens * num_topk num_m_blocks = ceil_div(capacity, block_m) num_n_blocks = ceil_div(l1_n, block_n) - num_compute_tiles = num_experts_per_rank * num_m_blocks * num_n_blocks + num_expert_waves = num_experts_per_rank // num_experts_per_wave + max_tiles_per_expert_wave = num_experts_per_wave * num_m_blocks * num_n_blocks + max_tile_rounds_per_expert_wave = ceil_div(max_tiles_per_expert_wave, num_sms) num_k_blocks = hidden // block_k # CTA roles: warps 0-1 dispatch, warps 2-3 TMA, then WGMMA warpgroups. warp_size = 32 @@ -209,7 +288,9 @@ def main( route_counts: T.Tensor((num_ranks, num_experts), T.int32), recv_counts: T.Tensor((num_experts_per_rank,), T.int32), route_slots: T.Tensor((num_tokens, num_topk), T.int32), - arrivals: T.Tensor((num_experts_per_rank, num_m_blocks), T.uint32), + arrivals: T.Tensor( + (num_experts_per_rank, ceil_div(capacity, block_m)), T.uint32 + ), recv_x: T.Tensor((num_experts_per_rank, capacity, hidden), T.float8_e4m3fn), recv_x_sf: T.Tensor((num_experts_per_rank, capacity, num_scale_groups), T.float32), recv_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), @@ -403,13 +484,55 @@ def main( if tx >= producer_begin and tx < producer_end: producer_step = T.alloc_var(T.int32, init=0) - for producer_wave in T.serial(ceil_div(num_compute_tiles, num_sms)): - producer_tile = bid + producer_wave * num_sms - if producer_tile < num_compute_tiles: - producer_n = producer_tile % num_n_blocks - producer_m = (producer_tile // num_n_blocks) % num_m_blocks - producer_expert = producer_tile // (num_n_blocks * num_m_blocks) - if producer_n * block_n < l1_n and producer_m * block_m < recv_counts[producer_expert]: + producer_wave_tile = T.alloc_var(T.int32, init=bid) + for producer_schedule_step in T.serial( + num_expert_waves * max_tile_rounds_per_expert_wave + ): + producer_expert_wave = ( + producer_schedule_step // max_tile_rounds_per_expert_wave + ) + producer_tile_round = ( + producer_schedule_step % max_tile_rounds_per_expert_wave + ) + producer_wave_begin = producer_expert_wave * num_experts_per_wave + producer_wave_num_tiles = T.alloc_var(T.int32, init=0) + for producer_wave_expert_offset in T.serial(num_experts_per_wave): + producer_wave_expert = ( + producer_wave_begin + producer_wave_expert_offset + ) + producer_wave_num_tiles += ( + T.ceildiv(T.min(recv_counts[producer_wave_expert], capacity), block_m) + * num_n_blocks + ) + + if producer_wave_tile < producer_wave_num_tiles: + producer_tile_offset = T.alloc_var( + T.int32, init=producer_wave_tile + ) + producer_expert = T.alloc_var(T.int32, init=-1) + producer_m = T.alloc_var(T.int32, init=0) + producer_n = T.alloc_var(T.int32, init=0) + for producer_wave_expert_offset in T.serial( + num_experts_per_wave + ): + producer_wave_expert = ( + producer_wave_begin + producer_wave_expert_offset + ) + producer_expert_m_blocks = T.ceildiv( + T.min(recv_counts[producer_wave_expert], capacity), block_m + ) + producer_expert_tiles = ( + producer_expert_m_blocks * num_n_blocks + ) + if producer_expert < 0: + if producer_tile_offset < producer_expert_tiles: + producer_expert = producer_wave_expert + producer_m = producer_tile_offset // num_n_blocks + producer_n = producer_tile_offset % num_n_blocks + else: + producer_tile_offset -= producer_expert_tiles + + if producer_expert >= 0 and producer_n * block_n < l1_n: producer_arrivals = T.min( block_m, recv_counts[producer_expert] - producer_m * block_m, @@ -457,6 +580,13 @@ def main( ) T.mbarrier_arrive(stage_barriers[producer_stage]) producer_step += num_k_blocks + producer_wave_tile += num_sms + + if ( + producer_tile_round == max_tile_rounds_per_expert_wave - 1 + and producer_wave_num_tiles > 0 + ): + producer_wave_tile -= producer_wave_num_tiles elif tx >= math_begin: partial = T.alloc_fragment((block_m, block_n), T.float32) @@ -473,14 +603,56 @@ def main( act_scale = T.alloc_fragment((block_m,), T.float32) weight_scale = T.alloc_local((2 * num_output_scale_groups,), T.float32) consumer_step = T.alloc_var(T.int32, init=0) + consumer_wave_tile = T.alloc_var(T.int32, init=bid) - for consumer_wave in T.serial(ceil_div(num_compute_tiles, num_sms)): - consumer_tile = bid + consumer_wave * num_sms - if consumer_tile < num_compute_tiles: - consumer_n = consumer_tile % num_n_blocks - consumer_m = (consumer_tile // num_n_blocks) % num_m_blocks - consumer_expert = consumer_tile // (num_n_blocks * num_m_blocks) - if consumer_n * block_n < l1_n and consumer_m * block_m < recv_counts[consumer_expert]: + for consumer_schedule_step in T.serial( + num_expert_waves * max_tile_rounds_per_expert_wave + ): + consumer_expert_wave = ( + consumer_schedule_step // max_tile_rounds_per_expert_wave + ) + consumer_tile_round = ( + consumer_schedule_step % max_tile_rounds_per_expert_wave + ) + consumer_wave_begin = consumer_expert_wave * num_experts_per_wave + consumer_wave_num_tiles = T.alloc_var(T.int32, init=0) + for consumer_wave_expert_offset in T.serial(num_experts_per_wave): + consumer_wave_expert = ( + consumer_wave_begin + consumer_wave_expert_offset + ) + consumer_wave_num_tiles += ( + T.ceildiv(T.min(recv_counts[consumer_wave_expert], capacity), block_m) + * num_n_blocks + ) + + if consumer_wave_tile < consumer_wave_num_tiles: + consumer_tile_offset = T.alloc_var( + T.int32, init=consumer_wave_tile + ) + consumer_expert = T.alloc_var(T.int32, init=-1) + consumer_m = T.alloc_var(T.int32, init=0) + consumer_n = T.alloc_var(T.int32, init=0) + for consumer_wave_expert_offset in T.serial( + num_experts_per_wave + ): + consumer_wave_expert = ( + consumer_wave_begin + consumer_wave_expert_offset + ) + consumer_expert_m_blocks = T.ceildiv( + T.min(recv_counts[consumer_wave_expert], capacity), block_m + ) + consumer_expert_tiles = ( + consumer_expert_m_blocks * num_n_blocks + ) + if consumer_expert < 0: + if consumer_tile_offset < consumer_expert_tiles: + consumer_expert = consumer_wave_expert + consumer_m = consumer_tile_offset // num_n_blocks + consumer_n = consumer_tile_offset % num_n_blocks + else: + consumer_tile_offset -= consumer_expert_tiles + + if consumer_expert >= 0 and consumer_n * block_n < l1_n: T.clear(partial) T.clear(accum) for consumer_k in T.serial(num_k_blocks): @@ -575,6 +747,13 @@ def main( consumer_n * (block_n // 2), ], ) + consumer_wave_tile += num_sms + + if ( + consumer_tile_round == max_tile_rounds_per_expert_wave - 1 + and consumer_wave_num_tiles > 0 + ): + consumer_wave_tile -= consumer_wave_num_tiles return main @@ -596,10 +775,15 @@ def fused_l2_scatter_reduce_manual_warp_kernel( reduce_block_h: int = 128, threads: int = 384, pipeline_stages: int = 3, + num_experts_per_wave: int | None = None, ): + num_experts_per_wave = num_experts_per_wave or num_experts_per_rank + assert num_experts_per_rank % num_experts_per_wave == 0 num_m_blocks = ceil_div(capacity, block_m) num_n_blocks = ceil_div(hidden, block_n) - num_compute_tiles = num_experts_per_rank * num_m_blocks * num_n_blocks + num_expert_waves = num_experts_per_rank // num_experts_per_wave + max_tiles_per_expert_wave = num_experts_per_wave * num_m_blocks * num_n_blocks + max_tile_rounds_per_expert_wave = ceil_div(max_tiles_per_expert_wave, num_sms) num_k_blocks = intermediate_hidden // block_k num_reduce_n_blocks = ceil_div(hidden, reduce_block_h) num_reduce_m_blocks = ceil_div(num_tokens, reduce_block_m) @@ -683,17 +867,58 @@ def main( if tx >= producer_begin and tx < producer_end: producer_step = T.alloc_var(T.int32, init=0) - for producer_wave in T.serial(ceil_div(num_compute_tiles, num_sms)): - producer_tile = bid + producer_wave * num_sms - if producer_tile < num_compute_tiles: - producer_n = producer_tile % num_n_blocks - producer_m = (producer_tile // num_n_blocks) % num_m_blocks - producer_expert = producer_tile // (num_n_blocks * num_m_blocks) + producer_wave_tile = T.alloc_var(T.int32, init=bid) + for producer_schedule_step in T.serial( + num_expert_waves * max_tile_rounds_per_expert_wave + ): + producer_expert_wave = ( + producer_schedule_step // max_tile_rounds_per_expert_wave + ) + producer_tile_round = ( + producer_schedule_step % max_tile_rounds_per_expert_wave + ) + producer_wave_begin = producer_expert_wave * num_experts_per_wave + producer_wave_num_tiles = T.alloc_var(T.int32, init=0) + for producer_wave_expert_offset in T.serial(num_experts_per_wave): + producer_wave_expert = ( + producer_wave_begin + producer_wave_expert_offset + ) + producer_wave_num_tiles += ( + T.ceildiv(T.min(recv_counts[producer_wave_expert], capacity), block_m) + * num_n_blocks + ) + + if producer_wave_tile < producer_wave_num_tiles: + producer_tile_offset = T.alloc_var( + T.int32, init=producer_wave_tile + ) + producer_expert = T.alloc_var(T.int32, init=-1) + producer_m = T.alloc_var(T.int32, init=0) + producer_n = T.alloc_var(T.int32, init=0) + for producer_wave_expert_offset in T.serial( + num_experts_per_wave + ): + producer_wave_expert = ( + producer_wave_begin + producer_wave_expert_offset + ) + producer_expert_m_blocks = T.ceildiv( + T.min(recv_counts[producer_wave_expert], capacity), block_m + ) + producer_expert_tiles = ( + producer_expert_m_blocks * num_n_blocks + ) + if producer_expert < 0: + if producer_tile_offset < producer_expert_tiles: + producer_expert = producer_wave_expert + producer_m = producer_tile_offset // num_n_blocks + producer_n = producer_tile_offset % num_n_blocks + else: + producer_tile_offset -= producer_expert_tiles + if ( - producer_expert < num_experts_per_rank + producer_expert >= 0 + and producer_expert < num_experts_per_rank and producer_n * block_n < hidden - and producer_m * block_m < capacity - and producer_m * block_m < recv_counts[producer_expert] and num_k_blocks * block_k == intermediate_hidden ): for producer_k in T.serial(num_k_blocks): @@ -728,6 +953,13 @@ def main( ] T.mbarrier_arrive(stage_barriers[producer_stage]) producer_step += num_k_blocks + producer_wave_tile += num_sms + + if ( + producer_tile_round == max_tile_rounds_per_expert_wave - 1 + and producer_wave_num_tiles > 0 + ): + producer_wave_tile -= producer_wave_num_tiles elif tx >= math_begin: partial = T.alloc_fragment((block_m, block_n), T.float32) @@ -735,21 +967,62 @@ def main( act_scale = T.alloc_fragment((block_m,), T.float32) weight_scale = T.alloc_local((2,), T.float32) consumer_step = T.alloc_var(T.int32, init=0) + consumer_wave_tile = T.alloc_var(T.int32, init=bid) scatter_dst_rank = T.alloc_var(T.int32, init=0) scatter_dst_token = T.alloc_var(T.int32, init=0) scatter_dst_topk = T.alloc_var(T.int32, init=0) - for consumer_wave in T.serial(ceil_div(num_compute_tiles, num_sms)): - consumer_tile = bid + consumer_wave * num_sms - if consumer_tile < num_compute_tiles: - consumer_n = consumer_tile % num_n_blocks - consumer_m = (consumer_tile // num_n_blocks) % num_m_blocks - consumer_expert = consumer_tile // (num_n_blocks * num_m_blocks) + for consumer_schedule_step in T.serial( + num_expert_waves * max_tile_rounds_per_expert_wave + ): + consumer_expert_wave = ( + consumer_schedule_step // max_tile_rounds_per_expert_wave + ) + consumer_tile_round = ( + consumer_schedule_step % max_tile_rounds_per_expert_wave + ) + consumer_wave_begin = consumer_expert_wave * num_experts_per_wave + consumer_wave_num_tiles = T.alloc_var(T.int32, init=0) + for consumer_wave_expert_offset in T.serial(num_experts_per_wave): + consumer_wave_expert = ( + consumer_wave_begin + consumer_wave_expert_offset + ) + consumer_wave_num_tiles += ( + T.ceildiv(T.min(recv_counts[consumer_wave_expert], capacity), block_m) + * num_n_blocks + ) + + if consumer_wave_tile < consumer_wave_num_tiles: + consumer_tile_offset = T.alloc_var( + T.int32, init=consumer_wave_tile + ) + consumer_expert = T.alloc_var(T.int32, init=-1) + consumer_m = T.alloc_var(T.int32, init=0) + consumer_n = T.alloc_var(T.int32, init=0) + for consumer_wave_expert_offset in T.serial( + num_experts_per_wave + ): + consumer_wave_expert = ( + consumer_wave_begin + consumer_wave_expert_offset + ) + consumer_expert_m_blocks = T.ceildiv( + T.min(recv_counts[consumer_wave_expert], capacity), block_m + ) + consumer_expert_tiles = ( + consumer_expert_m_blocks * num_n_blocks + ) + if consumer_expert < 0: + if consumer_tile_offset < consumer_expert_tiles: + consumer_expert = consumer_wave_expert + consumer_m = consumer_tile_offset // num_n_blocks + consumer_n = consumer_tile_offset % num_n_blocks + else: + consumer_tile_offset -= consumer_expert_tiles + if ( - consumer_expert < num_experts_per_rank + consumer_expert >= 0 + and consumer_expert < num_experts_per_rank and consumer_n * block_n < hidden - and consumer_m * block_m < capacity - and consumer_m * block_m < recv_counts[consumer_expert] and num_k_blocks * block_k == intermediate_hidden ): T.clear(partial) @@ -825,6 +1098,13 @@ def main( dst_pe=dst_rank, unroll_factor=1, ) + consumer_wave_tile += num_sms + + if ( + consumer_tile_round == max_tile_rounds_per_expert_wave - 1 + and consumer_wave_num_tiles > 0 + ): + consumer_wave_tile -= consumer_wave_num_tiles T.fence_sys() T.sync_grid() @@ -1006,6 +1286,11 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): num_experts_per_rank, num_sms, ) + for phase, config in (("l1", l1_config), ("l2", l2_config)): + requested = getattr(args, f"{phase}_experts_per_wave", None) + if requested is not None: + assert requested > 0 and num_experts_per_rank % requested == 0 + config["num_experts_per_wave"] = requested kernel_specs = [ fused_l1_swiglu_manual_warp_kernel( num_tokens, @@ -1201,7 +1486,9 @@ def run_pipeline(check_capacity: bool = False): print( f"tilescale sm90 fp8 mega moe: model={model_name} family={shape_family} " f"M={num_tokens} H={hidden} IH={intermediate_hidden} E={num_experts} " - f"topk={num_topk} capacity={capacity} latency={latency * 1000:.1f} us" + f"topk={num_topk} capacity={capacity} " + f"epw={l1_config['num_experts_per_wave']}/{l2_config['num_experts_per_wave']} " + f"latency={latency * 1000:.1f} us" ) allocator.close() @@ -1218,6 +1505,8 @@ def run_pipeline(check_capacity: bool = False): parser.add_argument("--num-topk", type=int, default=None) parser.add_argument("--num-tokens", type=int, default=64) parser.add_argument("--capacity", type=int, default=None) + parser.add_argument("--l1-experts-per-wave", type=int, default=None) + parser.add_argument("--l2-experts-per-wave", type=int, default=None) parser.add_argument("--activation-clamp", type=float, default=10.0) parser.add_argument("--seed", type=int, default=0) parser.add_argument("--diff-tol", type=float, default=0.01) diff --git a/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py index d73cb7618d..02b7cbbc8b 100644 --- a/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py @@ -40,6 +40,7 @@ def test_custom_model_config_and_schedule(): "block_k": 128, "threads": 384, "pipeline_stages": 5, + "num_experts_per_wave": 16, } assert l2 == {**l1, "pipeline_stages": 3} @@ -48,6 +49,7 @@ def test_custom_model_config_and_schedule(): ) assert family == "compact" assert l1["pipeline_stages"] == l2["pipeline_stages"] == 3 + assert l1["num_experts_per_wave"] == l2["num_experts_per_wave"] == 4 family, l1, l2 = example_sm90_fp8_mega_moe.select_manual_warp_configs( 7168, 3072, num_tokens=128, num_topk=6, num_experts_per_rank=48, num_sms=132 @@ -55,6 +57,8 @@ def test_custom_model_config_and_schedule(): assert family == "wide" assert l1["pipeline_stages"] == 4 assert l2["pipeline_stages"] == 3 + assert l1["num_experts_per_wave"] == l2["num_experts_per_wave"] == 4 + assert example_sm90_fp8_mega_moe.normalize_experts_per_wave(50, 16) == 25 @distributed_test(nprocs=4, require_fabric=True) From 4be976157e13cf8b029f1968ea6a3912f6d90867 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Fri, 14 Aug 2026 11:45:58 +0800 Subject: [PATCH 18/30] feat(distributed): trace mega MoE kernel pipeline --- .../mega_moe/example_sm90_fp8_mega_moe.py | 210 ++++++++++++++++++ .../distributed/mega_moe/pipeline_trace.py | 194 ++++++++++++++++ 2 files changed, 404 insertions(+) create mode 100644 examples/distributed/mega_moe/pipeline_trace.py diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 65aa7695b7..81d31e6e14 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -23,6 +23,15 @@ from tilelang.distributed.bench import do_bench from tilelang.distributed.host import init_dist +from pipeline_trace import ( + GLOBAL_TIMER_SOURCE, + L1_TRACE_ROLES, + L2_TRACE_ROLES, + TRACE_FIELDS, + save_pipeline_trace_png, + trace_schedule_steps, +) + os.environ.setdefault("NCCL_DEBUG", "ERROR") @@ -242,6 +251,7 @@ def fused_l1_swiglu_manual_warp_kernel( threads: int = 384, pipeline_stages: int = 5, num_experts_per_wave: int | None = None, + enable_pipeline_trace: bool = False, ): num_experts_per_rank = num_experts // num_ranks num_experts_per_wave = num_experts_per_wave or num_experts_per_rank @@ -308,9 +318,19 @@ def main( T.float32, ), barrier: T.Tensor((num_ranks,), T.int32), + pipeline_trace: T.Tensor( + ( + num_sms, + L1_TRACE_ROLES, + num_expert_waves * max_tile_rounds_per_expert_wave, + TRACE_FIELDS, + ), + T.int64, + ), ): T.annotate_pass_configs({tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) with T.Kernel(num_sms, threads=threads) as bid: + T.import_source(GLOBAL_TIMER_SOURCE) tx = T.get_thread_binding() src_rank = T.alloc_local((1,), T.int32) src_rank[0] = T.get_rank() @@ -403,6 +423,11 @@ def main( dispatch_warp = tx // warp_size dispatch_lane = tx % warp_size if tx < dispatch_threads: + if enable_pipeline_trace: + if dispatch_lane == 0: + pipeline_trace[bid, dispatch_warp, 0, 0] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) for metadata_wave in T.serial(ceil_div(num_routes, num_sms * dispatch_threads)): metadata_route = ( bid * dispatch_threads @@ -440,6 +465,15 @@ def main( dst_pe=metadata_rank, ) + if enable_pipeline_trace: + if dispatch_lane == 0: + pipeline_trace[bid, dispatch_warp, 0, 1] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) + pipeline_trace[bid, dispatch_warp, 0, 2] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) + for pull_wave in T.serial(ceil_div(num_experts_per_rank * capacity, num_sms * dispatch_warps)): pull_idx = ( bid * dispatch_warps @@ -482,6 +516,12 @@ def main( sem="release", ) + if enable_pipeline_trace: + if dispatch_lane == 0: + pipeline_trace[bid, dispatch_warp, 0, 3] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) + if tx >= producer_begin and tx < producer_end: producer_step = T.alloc_var(T.int32, init=0) producer_wave_tile = T.alloc_var(T.int32, init=bid) @@ -533,6 +573,11 @@ def main( producer_tile_offset -= producer_expert_tiles if producer_expert >= 0 and producer_n * block_n < l1_n: + if enable_pipeline_trace: + if tx == producer_begin: + pipeline_trace[bid, 2, producer_schedule_step, 0] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) producer_arrivals = T.min( block_m, recv_counts[producer_expert] - producer_m * block_m, @@ -545,6 +590,11 @@ def main( semantics=T.WaitSemantics.ACQUIRE, ) T.sync_threads(5, producer_threads) + if enable_pipeline_trace: + if tx == producer_begin: + pipeline_trace[bid, 2, producer_schedule_step, 1] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) for producer_k in T.serial(num_k_blocks): producer_stage = (producer_step + producer_k) % pipeline_stages producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 @@ -579,6 +629,11 @@ def main( barrier=stage_barriers[producer_stage], ) T.mbarrier_arrive(stage_barriers[producer_stage]) + if enable_pipeline_trace: + if tx == producer_begin: + pipeline_trace[bid, 2, producer_schedule_step, 2] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) producer_step += num_k_blocks producer_wave_tile += num_sms @@ -653,6 +708,11 @@ def main( consumer_tile_offset -= consumer_expert_tiles if consumer_expert >= 0 and consumer_n * block_n < l1_n: + if enable_pipeline_trace: + if tx == math_begin: + pipeline_trace[bid, 3, consumer_schedule_step, 3] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) T.clear(partial) T.clear(accum) for consumer_k in T.serial(num_k_blocks): @@ -662,6 +722,12 @@ def main( stage_barriers[consumer_stage], consumer_phase, ) + if consumer_k == 0: + if enable_pipeline_trace: + if tx == math_begin: + pipeline_trace[bid, 3, consumer_schedule_step, 0] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) T.gemm( a_shared[consumer_stage, :, :], b_shared[consumer_stage, :, :], @@ -703,6 +769,11 @@ def main( T.clear(partial) T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) consumer_step += num_k_blocks + if enable_pipeline_trace: + if tx == math_begin: + pipeline_trace[bid, 3, consumer_schedule_step, 1] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) for i, j in T.Parallel(block_m, block_n // 2): gate[i, j] = accum[i, (j // 8) * 16 + j % 8] for i, j in T.Parallel(block_m, block_n // 2): @@ -747,6 +818,11 @@ def main( consumer_n * (block_n // 2), ], ) + if enable_pipeline_trace: + if tx == math_begin: + pipeline_trace[bid, 3, consumer_schedule_step, 2] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) consumer_wave_tile += num_sms if ( @@ -776,6 +852,7 @@ def fused_l2_scatter_reduce_manual_warp_kernel( threads: int = 384, pipeline_stages: int = 3, num_experts_per_wave: int | None = None, + enable_pipeline_trace: bool = False, ): num_experts_per_wave = num_experts_per_wave or num_experts_per_rank assert num_experts_per_rank % num_experts_per_wave == 0 @@ -842,8 +919,18 @@ def main( combine: T.Tensor((num_tokens, num_topk, hidden), T.bfloat16), barrier: T.Tensor((num_ranks,), T.int32), out: T.Tensor((num_tokens, hidden), T.bfloat16), + pipeline_trace: T.Tensor( + ( + num_sms, + L2_TRACE_ROLES, + num_expert_waves * max_tile_rounds_per_expert_wave, + TRACE_FIELDS, + ), + T.int64, + ), ): with T.Kernel(num_sms, threads=threads) as bid: + T.import_source(GLOBAL_TIMER_SOURCE) tx = T.get_thread_binding() a_shared = T.alloc_shared( (pipeline_stages, block_m, block_k), @@ -921,6 +1008,11 @@ def main( and producer_n * block_n < hidden and num_k_blocks * block_k == intermediate_hidden ): + if enable_pipeline_trace: + if tx == producer_begin: + pipeline_trace[bid, 0, producer_schedule_step, 0] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) for producer_k in T.serial(num_k_blocks): producer_stage = (producer_step + producer_k) % pipeline_stages producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 @@ -952,6 +1044,11 @@ def main( producer_k, ] T.mbarrier_arrive(stage_barriers[producer_stage]) + if enable_pipeline_trace: + if tx == producer_begin: + pipeline_trace[bid, 0, producer_schedule_step, 1] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) producer_step += num_k_blocks producer_wave_tile += num_sms @@ -1025,6 +1122,11 @@ def main( and consumer_n * block_n < hidden and num_k_blocks * block_k == intermediate_hidden ): + if enable_pipeline_trace: + if tx == math_begin: + pipeline_trace[bid, 1, consumer_schedule_step, 3] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) T.clear(partial) T.clear(accum) for consumer_k in T.serial(num_k_blocks): @@ -1034,6 +1136,12 @@ def main( stage_barriers[consumer_stage], consumer_phase, ) + if consumer_k == 0: + if enable_pipeline_trace: + if tx == math_begin: + pipeline_trace[bid, 1, consumer_schedule_step, 0] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) T.gemm( a_shared[consumer_stage, :, :], b_shared[consumer_stage, :, :], @@ -1064,6 +1172,11 @@ def main( T.clear(partial) T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) consumer_step += num_k_blocks + if enable_pipeline_trace: + if tx == math_begin: + pipeline_trace[bid, 1, consumer_schedule_step, 1] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) T.copy(accum, out_shared) scatter_warp = (tx - math_begin) // warp_size for row_in_warp in T.serial(rows_per_math_warp): @@ -1098,6 +1211,11 @@ def main( dst_pe=dst_rank, unroll_factor=1, ) + if enable_pipeline_trace: + if tx == math_begin: + pipeline_trace[bid, 1, consumer_schedule_step, 2] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) consumer_wave_tile += num_sms if ( @@ -1114,6 +1232,11 @@ def main( if tx < reduce_threads: reduce_accum = T.alloc_fragment((reduce_block_m, reduce_block_h), T.float32) + if enable_pipeline_trace: + if tx == 0: + pipeline_trace[bid, 2, 0, 0] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) for reduce_wave in T.serial(ceil_div(num_reduce_tiles, num_sms)): reduce_tile = bid + reduce_wave * num_sms if reduce_tile < num_reduce_tiles: @@ -1136,6 +1259,11 @@ def main( reduce_n * reduce_block_h, ], ) + if enable_pipeline_trace: + if tx == 0: + pipeline_trace[bid, 2, 0, 1] = T.call_extern( + "int64", "tl_globaltimer_ns" + ) return main @@ -1243,6 +1371,8 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): num_topk = model["num_topk"] num_tokens = args.num_tokens activation_clamp = args.activation_clamp + pipeline_trace_path = getattr(args, "pipeline_trace", None) + enable_pipeline_trace = pipeline_trace_path is not None assert num_tokens > 0 assert hidden >= 512 and hidden % 256 == 0 @@ -1291,6 +1421,22 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): if requested is not None: assert requested > 0 and num_experts_per_rank % requested == 0 config["num_experts_per_wave"] = requested + l1_trace_steps = trace_schedule_steps( + num_experts_per_rank, + l1_config["num_experts_per_wave"], + capacity, + l1_config["block_m"], + ceil_div(2 * intermediate_hidden, l1_config["block_n"]), + num_sms, + ) + l2_trace_steps = trace_schedule_steps( + num_experts_per_rank, + l2_config["num_experts_per_wave"], + capacity, + l2_config["block_m"], + ceil_div(hidden, l2_config["block_n"]), + num_sms, + ) kernel_specs = [ fused_l1_swiglu_manual_warp_kernel( num_tokens, @@ -1302,6 +1448,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): capacity, num_sms, activation_clamp=activation_clamp, + enable_pipeline_trace=enable_pipeline_trace, **l1_config, ), fused_l2_scatter_reduce_manual_warp_kernel( @@ -1313,6 +1460,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): num_ranks, capacity, num_sms, + enable_pipeline_trace=enable_pipeline_trace, **l2_config, ), ] @@ -1392,6 +1540,16 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): ) combine = allocator_tensor((num_tokens, num_topk, hidden), torch.bfloat16, allocator=allocator) out = allocator_tensor((num_tokens, hidden), torch.bfloat16, allocator=allocator) + l1_pipeline_trace = torch.zeros( + (num_sms, L1_TRACE_ROLES, l1_trace_steps, TRACE_FIELDS), + dtype=torch.int64, + device="cuda", + ) + l2_pipeline_trace = torch.zeros( + (num_sms, L2_TRACE_ROLES, l2_trace_steps, TRACE_FIELDS), + dtype=torch.int64, + device="cuda", + ) def reset_state(): route_counts.zero_() @@ -1403,6 +1561,8 @@ def reset_state(): recv_weights.zero_() src_ranks.fill_(-1) combine.zero_() + l1_pipeline_trace.zero_() + l2_pipeline_trace.zero_() torch.cuda.synchronize() dist.barrier(group=group) @@ -1427,6 +1587,7 @@ def run_pipeline(check_capacity: bool = False): l2_x, l2_x_sf, barrier, + l1_pipeline_trace, ) if check_capacity: local_max = recv_counts.max() @@ -1444,6 +1605,7 @@ def run_pipeline(check_capacity: bool = False): combine, barrier, out, + l2_pipeline_trace, ) return out @@ -1469,6 +1631,52 @@ def run_pipeline(check_capacity: bool = False): assert diff < args.diff_tol, f"rank {local_rank}: diff={diff} exceeds {args.diff_tol}" print(f"rank {local_rank} check passed, diff={diff:.6f}") + if enable_pipeline_trace: + dist.barrier(group=group) + if local_rank == 0: + if getattr(args, "pipeline_trace_ctas", None): + selected_ctas = [ + int(value) + for value in args.pipeline_trace_ctas.split(",") + if value.strip() + ] + else: + local_counts = recv_counts.detach().cpu().tolist() + + def first_wave_tiles(config, output_size): + experts_per_wave = config["num_experts_per_wave"] + block_m = config["block_m"] + num_n_blocks = ceil_div(output_size, config["block_n"]) + return sum( + ceil_div(min(int(count), capacity), block_m) * num_n_blocks + for count in local_counts[:experts_per_wave] + ) + + selected_ctas = [0, num_sms - 1] + for wave_tiles in ( + first_wave_tiles(l1_config, 2 * intermediate_hidden), + first_wave_tiles(l2_config, hidden), + ): + if 0 < wave_tiles < num_sms: + selected_ctas.extend((wave_tiles - 1, wave_tiles)) + selected_ctas = sorted( + {cta for cta in selected_ctas if 0 <= cta < num_sms} + ) + trace_output = save_pipeline_trace_png( + l1_pipeline_trace, + l2_pipeline_trace, + pipeline_trace_path, + selected_ctas, + ( + f"TileScale SM90 Mega MoE pipeline: {model_name}, " + f"M={num_tokens}, ranks={num_ranks}, " + f"epw={l1_config['num_experts_per_wave']}/" + f"{l2_config['num_experts_per_wave']}" + ), + ) + print(f"pipeline trace: {trace_output} (CTAs={selected_ctas})") + dist.barrier(group=group) + if args.rep > 0: reset_state() # Stateful synchronization counters must be reset between warmup iterations. @@ -1507,6 +1715,8 @@ def run_pipeline(check_capacity: bool = False): parser.add_argument("--capacity", type=int, default=None) parser.add_argument("--l1-experts-per-wave", type=int, default=None) parser.add_argument("--l2-experts-per-wave", type=int, default=None) + parser.add_argument("--pipeline-trace", type=str, default=None) + parser.add_argument("--pipeline-trace-ctas", type=str, default=None) parser.add_argument("--activation-clamp", type=float, default=10.0) parser.add_argument("--seed", type=int, default=0) parser.add_argument("--diff-tol", type=float, default=0.01) diff --git a/examples/distributed/mega_moe/pipeline_trace.py b/examples/distributed/mega_moe/pipeline_trace.py new file mode 100644 index 0000000000..f8ec8ff340 --- /dev/null +++ b/examples/distributed/mega_moe/pipeline_trace.py @@ -0,0 +1,194 @@ +"""Low-overhead GPU timestamp tracing for the SM90 Mega MoE example.""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path +from typing import Iterable, Sequence + + +TRACE_FIELDS = 4 +L1_TRACE_ROLES = 4 +L2_TRACE_ROLES = 3 + +GLOBAL_TIMER_SOURCE = r""" +extern "C" __device__ __forceinline__ long long tl_globaltimer_ns() { + unsigned long long value; + asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(value)); + return static_cast(value); +} +""" + + +@dataclass(frozen=True) +class TraceRange: + cta: int + track: str + phase: str + start_ns: int + end_ns: int + + +PHASE_COLORS = { + "metadata": "#64748b", + "remote get": "#2563eb", + "arrival wait": "#eab308", + "stage wait": "#f97316", + "TMA": "#0d9488", + "GEMM": "#dc2626", + "epilogue": "#9333ea", + "scatter": "#0284c7", + "reduce": "#16a34a", +} + + +def trace_schedule_steps( + num_experts_per_rank: int, + num_experts_per_wave: int, + capacity: int, + block_m: int, + num_n_blocks: int, + num_sms: int, +) -> int: + num_waves = num_experts_per_rank // num_experts_per_wave + max_wave_tiles = num_experts_per_wave * ((capacity + block_m - 1) // block_m) * num_n_blocks + max_rounds = (max_wave_tiles + num_sms - 1) // num_sms + return num_waves * max_rounds + + +def _append_range( + ranges: list[TraceRange], + values: Sequence[int], + begin: int, + end: int, + cta: int, + track: str, + phase: str, +) -> None: + start_ns = int(values[begin]) + end_ns = int(values[end]) + if start_ns > 0 and end_ns >= start_ns: + ranges.append(TraceRange(cta, track, phase, start_ns, end_ns)) + + +def _collect_l1(trace, selected_ctas: Sequence[int]) -> list[TraceRange]: + values = trace.detach().cpu().tolist() + ranges: list[TraceRange] = [] + for cta in selected_ctas: + for dispatch_warp in range(2): + fields = values[cta][dispatch_warp][0] + track = f"dispatch{dispatch_warp}" + _append_range(ranges, fields, 0, 1, cta, track, "metadata") + _append_range(ranges, fields, 2, 3, cta, track, "remote get") + for fields in values[cta][2]: + _append_range(ranges, fields, 0, 1, cta, "producer", "arrival wait") + _append_range(ranges, fields, 1, 2, cta, "producer", "TMA") + for fields in values[cta][3]: + _append_range(ranges, fields, 3, 0, cta, "math", "stage wait") + _append_range(ranges, fields, 0, 1, cta, "math", "GEMM") + _append_range(ranges, fields, 1, 2, cta, "math", "epilogue") + return ranges + + +def _collect_l2(trace, selected_ctas: Sequence[int]) -> list[TraceRange]: + values = trace.detach().cpu().tolist() + ranges: list[TraceRange] = [] + for cta in selected_ctas: + for fields in values[cta][0]: + _append_range(ranges, fields, 0, 1, cta, "producer", "TMA") + for fields in values[cta][1]: + _append_range(ranges, fields, 3, 0, cta, "math", "stage wait") + _append_range(ranges, fields, 0, 1, cta, "math", "GEMM") + _append_range(ranges, fields, 1, 2, cta, "math", "scatter") + fields = values[cta][2][0] + _append_range(ranges, fields, 0, 1, cta, "reduce", "reduce") + return ranges + + +def _draw_panel(draw, ranges: Sequence[TraceRange], selected_ctas: Sequence[int], tracks: Sequence[str], box) -> None: + from PIL import ImageFont + + left, top, right, bottom = box + font = ImageFont.load_default() + label_width = 128 + axis_left = left + label_width + axis_right = right - 12 + title_height = 42 + row_height = max(16, (bottom - top - title_height - 24) // max(1, len(selected_ctas) * len(tracks))) + timestamps = [value for item in ranges for value in (item.start_ns, item.end_ns)] + if not timestamps: + draw.text((left + 8, top + 8), "No trace events", fill="#111827", font=font) + return + + start_ns = min(timestamps) + end_ns = max(timestamps) + duration_ns = max(1, end_ns - start_ns) + for tick in range(6): + x = axis_left + (axis_right - axis_left) * tick / 5 + draw.line((x, top + title_height, x, bottom - 16), fill="#e5e7eb", width=1) + label = f"{duration_ns * tick / 5000:.1f} us" + draw.text((x - 14, bottom - 14), label, fill="#475569", font=font) + + rows = [(cta, track) for cta in selected_ctas for track in tracks] + row_y = {(cta, track): top + title_height + index * row_height for index, (cta, track) in enumerate(rows)} + for index, (cta, track) in enumerate(rows): + y = row_y[(cta, track)] + if index % 2 == 0: + draw.rectangle((left, y, right, y + row_height - 1), fill="#f8fafc") + draw.text((left + 4, y + 2), f"CTA {cta:03d} {track}", fill="#1f2937", font=font) + + for item in ranges: + key = (item.cta, item.track) + if key not in row_y: + continue + x0 = axis_left + (item.start_ns - start_ns) * (axis_right - axis_left) / duration_ns + x1 = axis_left + (item.end_ns - start_ns) * (axis_right - axis_left) / duration_ns + y0 = row_y[key] + 2 + y1 = y0 + max(5, row_height - 5) + draw.rectangle((x0, y0, max(x0 + 1, x1), y1), fill=PHASE_COLORS[item.phase]) + + +def save_pipeline_trace_png( + l1_trace, + l2_trace, + output_path: str | Path, + selected_ctas: Iterable[int], + title: str, +) -> Path: + from PIL import Image, ImageDraw, ImageFont + + selected = tuple(dict.fromkeys(int(cta) for cta in selected_ctas)) + l1_ranges = _collect_l1(l1_trace, selected) + l2_ranges = _collect_l2(l2_trace, selected) + width = 1900 + rows = max(len(selected) * 4, len(selected) * 3) + height = max(420, 112 + rows * 22) + image = Image.new("RGB", (width, height), "white") + draw = ImageDraw.Draw(image) + font = ImageFont.load_default() + draw.text((18, 12), title, fill="#111827", font=font) + + legend_x = 18 + for phase, color in PHASE_COLORS.items(): + draw.rectangle((legend_x, 34, legend_x + 12, 46), fill=color) + draw.text((legend_x + 17, 34), phase, fill="#334155", font=font) + legend_x += 17 + 7 * len(phase) + 18 + + panel_top = 62 + panel_bottom = height - 8 + panel_mid = width // 2 + draw.text((18, panel_top), "L1: dispatch / TMA / GEMM / SwiGLU", fill="#111827", font=font) + draw.text((panel_mid + 8, panel_top), "L2: TMA / GEMM / scatter / reduce", fill="#111827", font=font) + _draw_panel( + draw, + l1_ranges, + selected, + ("dispatch0", "dispatch1", "producer", "math"), + (12, panel_top + 18, panel_mid - 4, panel_bottom), + ) + _draw_panel(draw, l2_ranges, selected, ("producer", "math", "reduce"), (panel_mid + 4, panel_top + 18, width - 12, panel_bottom)) + + output = Path(output_path).expanduser().absolute() + output.parent.mkdir(parents=True, exist_ok=True) + image.save(output) + return output From 00fe052847ac7f5d7e7eadc1db2131e7bb79878c Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Mon, 17 Aug 2026 15:14:12 +0800 Subject: [PATCH 19/30] refactor(distributed): streamline two-kernel mega MoE --- .../mega_moe/example_sm90_fp8_mega_moe.py | 545 +++++------------- .../distributed/mega_moe/pipeline_trace.py | 194 ------- 2 files changed, 150 insertions(+), 589 deletions(-) delete mode 100644 examples/distributed/mega_moe/pipeline_trace.py diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 81d31e6e14..4bc28a9864 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -23,15 +23,6 @@ from tilelang.distributed.bench import do_bench from tilelang.distributed.host import init_dist -from pipeline_trace import ( - GLOBAL_TIMER_SOURCE, - L1_TRACE_ROLES, - L2_TRACE_ROLES, - TRACE_FIELDS, - save_pipeline_trace_png, - trace_schedule_steps, -) - os.environ.setdefault("NCCL_DEBUG", "ERROR") @@ -45,14 +36,6 @@ SCALE_GRANULARITY = 128 -def ceil_div(x: int, y: int) -> int: - return (x + y - 1) // y - - -def align_up(x: int, alignment: int) -> int: - return ceil_div(x, alignment) * alignment - - def resolve_model_config(args: argparse.Namespace) -> Tuple[str, dict[str, int]]: model = MODEL_CONFIGS[args.model_config].copy() overrides = { @@ -66,14 +49,6 @@ def resolve_model_config(args: argparse.Namespace) -> Tuple[str, dict[str, int]] return ("custom" if is_custom else args.model_config), model -def classify_shape(hidden: int, intermediate_hidden: int) -> str: - if 3072 <= hidden < 5120 and 1536 <= intermediate_hidden < 2560: - return "compact" - if 5120 <= hidden <= 8192 and 2560 <= intermediate_hidden <= 4096: - return "wide" - return "generic" - - def normalize_experts_per_wave(num_experts: int, requested: int) -> int: requested = min(max(requested, 1), num_experts) for candidate in range(requested, num_experts + 1): @@ -82,38 +57,6 @@ def normalize_experts_per_wave(num_experts: int, requested: int) -> int: return num_experts -def select_generic_experts_per_wave( - intermediate_hidden: int, - routed_tokens: int, - num_experts_per_rank: int, - num_sms: int, - block_m: int = 64, - block_n: int = 256, -) -> int: - if routed_tokens < num_experts_per_rank or routed_tokens > 4 * num_experts_per_rank: - return num_experts_per_rank - - expected_tokens = ceil_div(routed_tokens, num_experts_per_rank) - num_m_blocks = ceil_div(expected_tokens, block_m) - num_n_blocks = 2 * intermediate_hidden // block_n - blocks_per_expert = num_m_blocks * num_n_blocks - requested = min( - num_experts_per_rank, - ceil_div(2 * num_sms, blocks_per_expert), - ) - if blocks_per_expert < num_sms: - max_candidate = min(num_experts_per_rank, 2 * requested) - requested = max( - range(requested, max_candidate + 1), - key=lambda candidate: ( - 1.0 - if num_experts_per_rank % candidate == 0 - else (num_experts_per_rank % candidate) / candidate - ), - ) - return normalize_experts_per_wave(num_experts_per_rank, requested) - - def select_manual_warp_configs( hidden: int, intermediate_hidden: int, @@ -123,18 +66,40 @@ def select_manual_warp_configs( num_sms: int, ) -> Tuple[str, dict[str, int], dict[str, int]]: """Select the TileScale counterpart of DeepGEMM SM90 schedule families.""" - shape_family = classify_shape(hidden, intermediate_hidden) + if 3072 <= hidden < 5120 and 1536 <= intermediate_hidden < 2560: + shape_family = "compact" + elif 5120 <= hidden <= 8192 and 2560 <= intermediate_hidden <= 4096: + shape_family = "wide" + else: + shape_family = "generic" + routed_tokens = num_tokens * num_topk high_sm = num_sms >= 100 l1_stages = 5 l2_stages = 3 - generic_experts_per_wave = select_generic_experts_per_wave( - intermediate_hidden, - routed_tokens, - num_experts_per_rank, - num_sms, - ) + generic_experts_per_wave = num_experts_per_rank + if num_experts_per_rank <= routed_tokens <= 4 * num_experts_per_rank: + expected_tokens = (routed_tokens + num_experts_per_rank - 1) // num_experts_per_rank + num_m_blocks = (expected_tokens + 63) // 64 + blocks_per_expert = num_m_blocks * (2 * intermediate_hidden // 256) + requested = min( + num_experts_per_rank, + (2 * num_sms + blocks_per_expert - 1) // blocks_per_expert, + ) + if blocks_per_expert < num_sms: + max_candidate = min(num_experts_per_rank, 2 * requested) + requested = max( + range(requested, max_candidate + 1), + key=lambda candidate: ( + 1.0 + if num_experts_per_rank % candidate == 0 + else (num_experts_per_rank % candidate) / candidate + ), + ) + generic_experts_per_wave = normalize_experts_per_wave( + num_experts_per_rank, requested + ) l1_experts_per_wave = l2_experts_per_wave = generic_experts_per_wave if high_sm and shape_family == "compact": if routed_tokens <= 32 * num_experts_per_rank: @@ -251,18 +216,17 @@ def fused_l1_swiglu_manual_warp_kernel( threads: int = 384, pipeline_stages: int = 5, num_experts_per_wave: int | None = None, - enable_pipeline_trace: bool = False, ): num_experts_per_rank = num_experts // num_ranks num_experts_per_wave = num_experts_per_wave or num_experts_per_rank assert num_experts_per_rank % num_experts_per_wave == 0 num_scale_groups = hidden // SCALE_GRANULARITY num_routes = num_tokens * num_topk - num_m_blocks = ceil_div(capacity, block_m) - num_n_blocks = ceil_div(l1_n, block_n) + num_m_blocks = T.ceildiv(capacity, block_m) + num_n_blocks = T.ceildiv(l1_n, block_n) num_expert_waves = num_experts_per_rank // num_experts_per_wave max_tiles_per_expert_wave = num_experts_per_wave * num_m_blocks * num_n_blocks - max_tile_rounds_per_expert_wave = ceil_div(max_tiles_per_expert_wave, num_sms) + max_tile_rounds_per_expert_wave = T.ceildiv(max_tiles_per_expert_wave, num_sms) num_k_blocks = hidden // block_k # CTA roles: warps 0-1 dispatch, warps 2-3 TMA, then WGMMA warpgroups. warp_size = 32 @@ -299,7 +263,7 @@ def main( recv_counts: T.Tensor((num_experts_per_rank,), T.int32), route_slots: T.Tensor((num_tokens, num_topk), T.int32), arrivals: T.Tensor( - (num_experts_per_rank, ceil_div(capacity, block_m)), T.uint32 + (num_experts_per_rank, T.ceildiv(capacity, block_m)), T.uint32 ), recv_x: T.Tensor((num_experts_per_rank, capacity, hidden), T.float8_e4m3fn), recv_x_sf: T.Tensor((num_experts_per_rank, capacity, num_scale_groups), T.float32), @@ -318,19 +282,9 @@ def main( T.float32, ), barrier: T.Tensor((num_ranks,), T.int32), - pipeline_trace: T.Tensor( - ( - num_sms, - L1_TRACE_ROLES, - num_expert_waves * max_tile_rounds_per_expert_wave, - TRACE_FIELDS, - ), - T.int64, - ), ): T.annotate_pass_configs({tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) with T.Kernel(num_sms, threads=threads) as bid: - T.import_source(GLOBAL_TIMER_SOURCE) tx = T.get_thread_binding() src_rank = T.alloc_local((1,), T.int32) src_rank[0] = T.get_rank() @@ -350,68 +304,62 @@ def main( if bid == 0: if tx < route_threads: - for reset_wave in T.serial(ceil_div(num_experts, route_threads)): - reset_expert = tx + reset_wave * route_threads - if reset_expert < num_experts: - route_counts[src_rank[0], reset_expert] = 0 + for reset_expert in T.serial(tx, num_experts, route_threads): + route_counts[src_rank[0], reset_expert] = 0 T.sync_threads(7, route_threads) - for assign_wave in T.serial(ceil_div(num_routes, route_threads)): - assign_route = tx + assign_wave * route_threads - if assign_route < num_routes: - assign_token = assign_route // num_topk - assign_topk = assign_route % num_topk - assign_expert = topk_idx[assign_token, assign_topk] - if assign_expert >= 0 and assign_expert < num_experts: - route_slots[assign_token, assign_topk] = T.atomic_add( - route_counts[src_rank[0], assign_expert], - 1, - memory_order="relaxed", - return_prev=True, - ) - else: - route_slots[assign_token, assign_topk] = -1 + for assign_route in T.serial(tx, num_routes, route_threads): + assign_token = assign_route // num_topk + assign_topk = assign_route % num_topk + assign_expert = topk_idx[assign_token, assign_topk] + if assign_expert >= 0 and assign_expert < num_experts: + route_slots[assign_token, assign_topk] = T.atomic_add( + route_counts[src_rank[0], assign_expert], + 1, + memory_order="relaxed", + return_prev=True, + ) + else: + route_slots[assign_token, assign_topk] = -1 T.sync_threads(7, route_threads) - for publish_wave in T.serial(ceil_div(num_experts * num_ranks, route_threads)): - publish_idx = tx + publish_wave * route_threads - if publish_idx < num_experts * num_ranks: - publish_rank = publish_idx // num_experts - publish_expert = publish_idx % num_experts - if publish_rank != src_rank[0]: - T.st( - route_counts[src_rank[0], publish_expert], - route_counts[src_rank[0], publish_expert], - dst_pe=publish_rank, - ) + for publish_idx in T.serial( + tx, num_experts * num_ranks, route_threads + ): + publish_rank = publish_idx // num_experts + publish_expert = publish_idx % num_experts + if publish_rank != src_rank[0]: + T.st( + route_counts[src_rank[0], publish_expert], + route_counts[src_rank[0], publish_expert], + dst_pe=publish_rank, + ) T.barrier_blocks(barrier[0]) if tx < route_threads: - for count_wave in T.serial(ceil_div(num_experts_per_rank, route_threads)): - count_local_expert = tx + count_wave * route_threads - if count_local_expert < num_experts_per_rank: - recv_count = T.alloc_var(T.int32, init=0) - recv_expert = src_rank[0] * num_experts_per_rank + count_local_expert - for count_rank in T.serial(num_ranks): - recv_count += route_counts[count_rank, recv_expert] - recv_counts[count_local_expert] = recv_count - - for prefix_wave in T.serial(ceil_div(num_routes, route_threads)): - prefix_route = tx + prefix_wave * route_threads - if prefix_route < num_routes: - prefix_token = prefix_route // num_topk - prefix_topk = prefix_route % num_topk - prefix_expert = topk_idx[prefix_token, prefix_topk] - prefix_slot = T.alloc_var( - T.int32, - init=route_slots[prefix_token, prefix_topk], - ) - if prefix_token < num_tokens and prefix_expert >= 0 and prefix_expert < num_experts and prefix_slot >= 0: - for prefix_rank in T.serial(num_ranks): - if prefix_rank < src_rank[0]: - prefix_slot += route_counts[prefix_rank, prefix_expert] - route_slots[prefix_token, prefix_topk] = prefix_slot + for count_local_expert in T.serial( + tx, num_experts_per_rank, route_threads + ): + recv_count = T.alloc_var(T.int32, init=0) + recv_expert = src_rank[0] * num_experts_per_rank + count_local_expert + for count_rank in T.serial(num_ranks): + recv_count += route_counts[count_rank, recv_expert] + recv_counts[count_local_expert] = recv_count + + for prefix_route in T.serial(tx, num_routes, route_threads): + prefix_token = prefix_route // num_topk + prefix_topk = prefix_route % num_topk + prefix_expert = topk_idx[prefix_token, prefix_topk] + prefix_slot = T.alloc_var( + T.int32, + init=route_slots[prefix_token, prefix_topk], + ) + if prefix_token < num_tokens and prefix_expert >= 0 and prefix_expert < num_experts and prefix_slot >= 0: + for prefix_rank in T.serial(num_ranks): + if prefix_rank < src_rank[0]: + prefix_slot += route_counts[prefix_rank, prefix_expert] + route_slots[prefix_token, prefix_topk] = prefix_slot T.sync_grid() @@ -423,66 +371,49 @@ def main( dispatch_warp = tx // warp_size dispatch_lane = tx % warp_size if tx < dispatch_threads: - if enable_pipeline_trace: - if dispatch_lane == 0: - pipeline_trace[bid, dispatch_warp, 0, 0] = T.call_extern( - "int64", "tl_globaltimer_ns" + for metadata_route in T.serial( + bid * dispatch_threads + tx, + num_routes, + num_sms * dispatch_threads, + ): + metadata_token = metadata_route // num_topk + metadata_topk = metadata_route % num_topk + metadata_expert = topk_idx[metadata_token, metadata_topk] + metadata_slot = route_slots[metadata_token, metadata_topk] + if metadata_expert >= 0 and metadata_slot >= 0 and metadata_slot < capacity: + metadata_rank = metadata_expert // num_experts_per_rank + metadata_local_expert = metadata_expert % num_experts_per_rank + T.st( + recv_weights[metadata_local_expert, metadata_slot], + topk_weights[metadata_token, metadata_topk], + dst_pe=metadata_rank, ) - for metadata_wave in T.serial(ceil_div(num_routes, num_sms * dispatch_threads)): - metadata_route = ( - bid * dispatch_threads - + tx - + metadata_wave * num_sms * dispatch_threads - ) - if metadata_route < num_routes: - metadata_token = metadata_route // num_topk - metadata_topk = metadata_route % num_topk - metadata_expert = topk_idx[metadata_token, metadata_topk] - metadata_slot = route_slots[metadata_token, metadata_topk] - if metadata_expert >= 0 and metadata_slot >= 0 and metadata_slot < capacity: - metadata_rank = metadata_expert // num_experts_per_rank - metadata_local_expert = metadata_expert % num_experts_per_rank - T.st( - recv_weights[metadata_local_expert, metadata_slot], - topk_weights[metadata_token, metadata_topk], - dst_pe=metadata_rank, - ) - T.st( - src_tokens[metadata_local_expert, metadata_slot], - metadata_token, - dst_pe=metadata_rank, - ) - T.st( - src_topk[metadata_local_expert, metadata_slot], - metadata_topk, - dst_pe=metadata_rank, - ) - T.st( - src_ranks[metadata_local_expert, metadata_slot], - src_rank[0], - scope="sys", - sem="release", - dst_pe=metadata_rank, - ) - - if enable_pipeline_trace: - if dispatch_lane == 0: - pipeline_trace[bid, dispatch_warp, 0, 1] = T.call_extern( - "int64", "tl_globaltimer_ns" + T.st( + src_tokens[metadata_local_expert, metadata_slot], + metadata_token, + dst_pe=metadata_rank, + ) + T.st( + src_topk[metadata_local_expert, metadata_slot], + metadata_topk, + dst_pe=metadata_rank, ) - pipeline_trace[bid, dispatch_warp, 0, 2] = T.call_extern( - "int64", "tl_globaltimer_ns" + T.st( + src_ranks[metadata_local_expert, metadata_slot], + src_rank[0], + scope="sys", + sem="release", + dst_pe=metadata_rank, ) - for pull_wave in T.serial(ceil_div(num_experts_per_rank * capacity, num_sms * dispatch_warps)): - pull_idx = ( - bid * dispatch_warps - + dispatch_warp - + pull_wave * num_sms * dispatch_warps - ) + for pull_idx in T.serial( + bid * dispatch_warps + dispatch_warp, + num_experts_per_rank * capacity, + num_sms * dispatch_warps, + ): pull_expert = pull_idx // capacity pull_slot = pull_idx % capacity - if pull_expert < num_experts_per_rank and pull_slot < recv_counts[pull_expert]: + if pull_slot < recv_counts[pull_expert]: if dispatch_lane == dispatch_leader_lane: T.wait_ge( src_ranks[pull_expert, pull_slot], @@ -516,12 +447,6 @@ def main( sem="release", ) - if enable_pipeline_trace: - if dispatch_lane == 0: - pipeline_trace[bid, dispatch_warp, 0, 3] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) - if tx >= producer_begin and tx < producer_end: producer_step = T.alloc_var(T.int32, init=0) producer_wave_tile = T.alloc_var(T.int32, init=bid) @@ -573,11 +498,6 @@ def main( producer_tile_offset -= producer_expert_tiles if producer_expert >= 0 and producer_n * block_n < l1_n: - if enable_pipeline_trace: - if tx == producer_begin: - pipeline_trace[bid, 2, producer_schedule_step, 0] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) producer_arrivals = T.min( block_m, recv_counts[producer_expert] - producer_m * block_m, @@ -590,11 +510,6 @@ def main( semantics=T.WaitSemantics.ACQUIRE, ) T.sync_threads(5, producer_threads) - if enable_pipeline_trace: - if tx == producer_begin: - pipeline_trace[bid, 2, producer_schedule_step, 1] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) for producer_k in T.serial(num_k_blocks): producer_stage = (producer_step + producer_k) % pipeline_stages producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 @@ -629,11 +544,6 @@ def main( barrier=stage_barriers[producer_stage], ) T.mbarrier_arrive(stage_barriers[producer_stage]) - if enable_pipeline_trace: - if tx == producer_begin: - pipeline_trace[bid, 2, producer_schedule_step, 2] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) producer_step += num_k_blocks producer_wave_tile += num_sms @@ -708,11 +618,6 @@ def main( consumer_tile_offset -= consumer_expert_tiles if consumer_expert >= 0 and consumer_n * block_n < l1_n: - if enable_pipeline_trace: - if tx == math_begin: - pipeline_trace[bid, 3, consumer_schedule_step, 3] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) T.clear(partial) T.clear(accum) for consumer_k in T.serial(num_k_blocks): @@ -722,12 +627,6 @@ def main( stage_barriers[consumer_stage], consumer_phase, ) - if consumer_k == 0: - if enable_pipeline_trace: - if tx == math_begin: - pipeline_trace[bid, 3, consumer_schedule_step, 0] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) T.gemm( a_shared[consumer_stage, :, :], b_shared[consumer_stage, :, :], @@ -769,11 +668,6 @@ def main( T.clear(partial) T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) consumer_step += num_k_blocks - if enable_pipeline_trace: - if tx == math_begin: - pipeline_trace[bid, 3, consumer_schedule_step, 1] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) for i, j in T.Parallel(block_m, block_n // 2): gate[i, j] = accum[i, (j // 8) * 16 + j % 8] for i, j in T.Parallel(block_m, block_n // 2): @@ -818,11 +712,6 @@ def main( consumer_n * (block_n // 2), ], ) - if enable_pipeline_trace: - if tx == math_begin: - pipeline_trace[bid, 3, consumer_schedule_step, 2] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) consumer_wave_tile += num_sms if ( @@ -852,18 +741,17 @@ def fused_l2_scatter_reduce_manual_warp_kernel( threads: int = 384, pipeline_stages: int = 3, num_experts_per_wave: int | None = None, - enable_pipeline_trace: bool = False, ): num_experts_per_wave = num_experts_per_wave or num_experts_per_rank assert num_experts_per_rank % num_experts_per_wave == 0 - num_m_blocks = ceil_div(capacity, block_m) - num_n_blocks = ceil_div(hidden, block_n) + num_m_blocks = T.ceildiv(capacity, block_m) + num_n_blocks = T.ceildiv(hidden, block_n) num_expert_waves = num_experts_per_rank // num_experts_per_wave max_tiles_per_expert_wave = num_experts_per_wave * num_m_blocks * num_n_blocks - max_tile_rounds_per_expert_wave = ceil_div(max_tiles_per_expert_wave, num_sms) + max_tile_rounds_per_expert_wave = T.ceildiv(max_tiles_per_expert_wave, num_sms) num_k_blocks = intermediate_hidden // block_k - num_reduce_n_blocks = ceil_div(hidden, reduce_block_h) - num_reduce_m_blocks = ceil_div(num_tokens, reduce_block_m) + num_reduce_n_blocks = T.ceildiv(hidden, reduce_block_h) + num_reduce_m_blocks = T.ceildiv(num_tokens, reduce_block_m) num_reduce_tiles = num_reduce_n_blocks * num_reduce_m_blocks # GEMM roles: warps 0-1 idle, warps 2-3 TMA, then two WGMMA # warpgroups. After the grid barrier, warps 0-3 perform top-k reduction. @@ -919,18 +807,8 @@ def main( combine: T.Tensor((num_tokens, num_topk, hidden), T.bfloat16), barrier: T.Tensor((num_ranks,), T.int32), out: T.Tensor((num_tokens, hidden), T.bfloat16), - pipeline_trace: T.Tensor( - ( - num_sms, - L2_TRACE_ROLES, - num_expert_waves * max_tile_rounds_per_expert_wave, - TRACE_FIELDS, - ), - T.int64, - ), ): with T.Kernel(num_sms, threads=threads) as bid: - T.import_source(GLOBAL_TIMER_SOURCE) tx = T.get_thread_binding() a_shared = T.alloc_shared( (pipeline_stages, block_m, block_k), @@ -1008,11 +886,6 @@ def main( and producer_n * block_n < hidden and num_k_blocks * block_k == intermediate_hidden ): - if enable_pipeline_trace: - if tx == producer_begin: - pipeline_trace[bid, 0, producer_schedule_step, 0] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) for producer_k in T.serial(num_k_blocks): producer_stage = (producer_step + producer_k) % pipeline_stages producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 @@ -1044,11 +917,6 @@ def main( producer_k, ] T.mbarrier_arrive(stage_barriers[producer_stage]) - if enable_pipeline_trace: - if tx == producer_begin: - pipeline_trace[bid, 0, producer_schedule_step, 1] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) producer_step += num_k_blocks producer_wave_tile += num_sms @@ -1122,11 +990,6 @@ def main( and consumer_n * block_n < hidden and num_k_blocks * block_k == intermediate_hidden ): - if enable_pipeline_trace: - if tx == math_begin: - pipeline_trace[bid, 1, consumer_schedule_step, 3] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) T.clear(partial) T.clear(accum) for consumer_k in T.serial(num_k_blocks): @@ -1136,12 +999,6 @@ def main( stage_barriers[consumer_stage], consumer_phase, ) - if consumer_k == 0: - if enable_pipeline_trace: - if tx == math_begin: - pipeline_trace[bid, 1, consumer_schedule_step, 0] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) T.gemm( a_shared[consumer_stage, :, :], b_shared[consumer_stage, :, :], @@ -1172,11 +1029,6 @@ def main( T.clear(partial) T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) consumer_step += num_k_blocks - if enable_pipeline_trace: - if tx == math_begin: - pipeline_trace[bid, 1, consumer_schedule_step, 1] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) T.copy(accum, out_shared) scatter_warp = (tx - math_begin) // warp_size for row_in_warp in T.serial(rows_per_math_warp): @@ -1211,11 +1063,6 @@ def main( dst_pe=dst_rank, unroll_factor=1, ) - if enable_pipeline_trace: - if tx == math_begin: - pipeline_trace[bid, 1, consumer_schedule_step, 2] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) consumer_wave_tile += num_sms if ( @@ -1232,39 +1079,26 @@ def main( if tx < reduce_threads: reduce_accum = T.alloc_fragment((reduce_block_m, reduce_block_h), T.float32) - if enable_pipeline_trace: - if tx == 0: - pipeline_trace[bid, 2, 0, 0] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) - for reduce_wave in T.serial(ceil_div(num_reduce_tiles, num_sms)): - reduce_tile = bid + reduce_wave * num_sms - if reduce_tile < num_reduce_tiles: - reduce_n = reduce_tile % num_reduce_n_blocks - reduce_m = reduce_tile // num_reduce_n_blocks - T.clear(reduce_accum) - for topk_slot in T.serial(num_topk): - for i, j in T.Parallel(reduce_block_m, reduce_block_h): - if reduce_m * reduce_block_m + i < num_tokens: - reduce_accum[i, j] += combine[ - reduce_m * reduce_block_m + i, - topk_slot, - reduce_n * reduce_block_h + j, - ] - T.copy(reduce_accum, reduce_shared) - T.copy( - reduce_shared, - out[ - reduce_m * reduce_block_m, - reduce_n * reduce_block_h, - ], - ) - if enable_pipeline_trace: - if tx == 0: - pipeline_trace[bid, 2, 0, 1] = T.call_extern( - "int64", "tl_globaltimer_ns" - ) - + for reduce_tile in T.serial(bid, num_reduce_tiles, num_sms): + reduce_n = reduce_tile % num_reduce_n_blocks + reduce_m = reduce_tile // num_reduce_n_blocks + T.clear(reduce_accum) + for topk_slot in T.serial(num_topk): + for i, j in T.Parallel(reduce_block_m, reduce_block_h): + if reduce_m * reduce_block_m + i < num_tokens: + reduce_accum[i, j] += combine[ + reduce_m * reduce_block_m + i, + topk_slot, + reduce_n * reduce_block_h + j, + ] + T.copy(reduce_accum, reduce_shared) + T.copy( + reduce_shared, + out[ + reduce_m * reduce_block_m, + reduce_n * reduce_block_h, + ], + ) return main @@ -1302,7 +1136,8 @@ def _allocator_size_bytes( + num_topk * (3 * i32 + fp32) + (num_topk + 1) * hidden * bf16 ) - return align_up(weight_bytes + weight_scale_bytes + pool_bytes + input_bytes + 2**27, 2**20) + total_bytes = weight_bytes + weight_scale_bytes + pool_bytes + input_bytes + 2**27 + return (total_bytes + 2**20 - 1) // 2**20 * 2**20 def _gather_cat(tensor: torch.Tensor, group: dist.ProcessGroup) -> torch.Tensor: @@ -1371,8 +1206,6 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): num_topk = model["num_topk"] num_tokens = args.num_tokens activation_clamp = args.activation_clamp - pipeline_trace_path = getattr(args, "pipeline_trace", None) - enable_pipeline_trace = pipeline_trace_path is not None assert num_tokens > 0 assert hidden >= 512 and hidden % 256 == 0 @@ -1380,11 +1213,13 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): assert num_experts > 0 and num_experts % num_local_ranks == 0 assert 0 < num_topk <= min(32, num_experts) num_experts_per_rank = num_experts // num_local_ranks - average_recv = ceil_div(num_tokens * num_local_ranks * num_topk, num_experts) + average_recv = ( + num_tokens * num_local_ranks * num_topk + num_experts - 1 + ) // num_experts capacity = ( args.capacity if args.capacity is not None - else align_up(max(average_recv * 2, 64), 64) + else (max(average_recv * 2, 64) + 63) // 64 * 64 ) assert capacity >= 64 and capacity % 64 == 0 @@ -1421,22 +1256,6 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): if requested is not None: assert requested > 0 and num_experts_per_rank % requested == 0 config["num_experts_per_wave"] = requested - l1_trace_steps = trace_schedule_steps( - num_experts_per_rank, - l1_config["num_experts_per_wave"], - capacity, - l1_config["block_m"], - ceil_div(2 * intermediate_hidden, l1_config["block_n"]), - num_sms, - ) - l2_trace_steps = trace_schedule_steps( - num_experts_per_rank, - l2_config["num_experts_per_wave"], - capacity, - l2_config["block_m"], - ceil_div(hidden, l2_config["block_n"]), - num_sms, - ) kernel_specs = [ fused_l1_swiglu_manual_warp_kernel( num_tokens, @@ -1448,7 +1267,6 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): capacity, num_sms, activation_clamp=activation_clamp, - enable_pipeline_trace=enable_pipeline_trace, **l1_config, ), fused_l2_scatter_reduce_manual_warp_kernel( @@ -1460,7 +1278,6 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): num_ranks, capacity, num_sms, - enable_pipeline_trace=enable_pipeline_trace, **l2_config, ), ] @@ -1528,7 +1345,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): src_topk = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) route_slots = allocator_tensor((num_tokens, num_topk), torch.int32, allocator=allocator) arrivals = allocator_tensor( - (num_experts_per_rank, ceil_div(capacity, 64)), + (num_experts_per_rank, (capacity + 63) // 64), torch.uint32, allocator=allocator, ) @@ -1540,16 +1357,6 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): ) combine = allocator_tensor((num_tokens, num_topk, hidden), torch.bfloat16, allocator=allocator) out = allocator_tensor((num_tokens, hidden), torch.bfloat16, allocator=allocator) - l1_pipeline_trace = torch.zeros( - (num_sms, L1_TRACE_ROLES, l1_trace_steps, TRACE_FIELDS), - dtype=torch.int64, - device="cuda", - ) - l2_pipeline_trace = torch.zeros( - (num_sms, L2_TRACE_ROLES, l2_trace_steps, TRACE_FIELDS), - dtype=torch.int64, - device="cuda", - ) def reset_state(): route_counts.zero_() @@ -1561,8 +1368,6 @@ def reset_state(): recv_weights.zero_() src_ranks.fill_(-1) combine.zero_() - l1_pipeline_trace.zero_() - l2_pipeline_trace.zero_() torch.cuda.synchronize() dist.barrier(group=group) @@ -1587,7 +1392,6 @@ def run_pipeline(check_capacity: bool = False): l2_x, l2_x_sf, barrier, - l1_pipeline_trace, ) if check_capacity: local_max = recv_counts.max() @@ -1605,7 +1409,6 @@ def run_pipeline(check_capacity: bool = False): combine, barrier, out, - l2_pipeline_trace, ) return out @@ -1631,52 +1434,6 @@ def run_pipeline(check_capacity: bool = False): assert diff < args.diff_tol, f"rank {local_rank}: diff={diff} exceeds {args.diff_tol}" print(f"rank {local_rank} check passed, diff={diff:.6f}") - if enable_pipeline_trace: - dist.barrier(group=group) - if local_rank == 0: - if getattr(args, "pipeline_trace_ctas", None): - selected_ctas = [ - int(value) - for value in args.pipeline_trace_ctas.split(",") - if value.strip() - ] - else: - local_counts = recv_counts.detach().cpu().tolist() - - def first_wave_tiles(config, output_size): - experts_per_wave = config["num_experts_per_wave"] - block_m = config["block_m"] - num_n_blocks = ceil_div(output_size, config["block_n"]) - return sum( - ceil_div(min(int(count), capacity), block_m) * num_n_blocks - for count in local_counts[:experts_per_wave] - ) - - selected_ctas = [0, num_sms - 1] - for wave_tiles in ( - first_wave_tiles(l1_config, 2 * intermediate_hidden), - first_wave_tiles(l2_config, hidden), - ): - if 0 < wave_tiles < num_sms: - selected_ctas.extend((wave_tiles - 1, wave_tiles)) - selected_ctas = sorted( - {cta for cta in selected_ctas if 0 <= cta < num_sms} - ) - trace_output = save_pipeline_trace_png( - l1_pipeline_trace, - l2_pipeline_trace, - pipeline_trace_path, - selected_ctas, - ( - f"TileScale SM90 Mega MoE pipeline: {model_name}, " - f"M={num_tokens}, ranks={num_ranks}, " - f"epw={l1_config['num_experts_per_wave']}/" - f"{l2_config['num_experts_per_wave']}" - ), - ) - print(f"pipeline trace: {trace_output} (CTAs={selected_ctas})") - dist.barrier(group=group) - if args.rep > 0: reset_state() # Stateful synchronization counters must be reset between warmup iterations. @@ -1715,8 +1472,6 @@ def first_wave_tiles(config, output_size): parser.add_argument("--capacity", type=int, default=None) parser.add_argument("--l1-experts-per-wave", type=int, default=None) parser.add_argument("--l2-experts-per-wave", type=int, default=None) - parser.add_argument("--pipeline-trace", type=str, default=None) - parser.add_argument("--pipeline-trace-ctas", type=str, default=None) parser.add_argument("--activation-clamp", type=float, default=10.0) parser.add_argument("--seed", type=int, default=0) parser.add_argument("--diff-tol", type=float, default=0.01) diff --git a/examples/distributed/mega_moe/pipeline_trace.py b/examples/distributed/mega_moe/pipeline_trace.py deleted file mode 100644 index f8ec8ff340..0000000000 --- a/examples/distributed/mega_moe/pipeline_trace.py +++ /dev/null @@ -1,194 +0,0 @@ -"""Low-overhead GPU timestamp tracing for the SM90 Mega MoE example.""" - -from __future__ import annotations - -from dataclasses import dataclass -from pathlib import Path -from typing import Iterable, Sequence - - -TRACE_FIELDS = 4 -L1_TRACE_ROLES = 4 -L2_TRACE_ROLES = 3 - -GLOBAL_TIMER_SOURCE = r""" -extern "C" __device__ __forceinline__ long long tl_globaltimer_ns() { - unsigned long long value; - asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(value)); - return static_cast(value); -} -""" - - -@dataclass(frozen=True) -class TraceRange: - cta: int - track: str - phase: str - start_ns: int - end_ns: int - - -PHASE_COLORS = { - "metadata": "#64748b", - "remote get": "#2563eb", - "arrival wait": "#eab308", - "stage wait": "#f97316", - "TMA": "#0d9488", - "GEMM": "#dc2626", - "epilogue": "#9333ea", - "scatter": "#0284c7", - "reduce": "#16a34a", -} - - -def trace_schedule_steps( - num_experts_per_rank: int, - num_experts_per_wave: int, - capacity: int, - block_m: int, - num_n_blocks: int, - num_sms: int, -) -> int: - num_waves = num_experts_per_rank // num_experts_per_wave - max_wave_tiles = num_experts_per_wave * ((capacity + block_m - 1) // block_m) * num_n_blocks - max_rounds = (max_wave_tiles + num_sms - 1) // num_sms - return num_waves * max_rounds - - -def _append_range( - ranges: list[TraceRange], - values: Sequence[int], - begin: int, - end: int, - cta: int, - track: str, - phase: str, -) -> None: - start_ns = int(values[begin]) - end_ns = int(values[end]) - if start_ns > 0 and end_ns >= start_ns: - ranges.append(TraceRange(cta, track, phase, start_ns, end_ns)) - - -def _collect_l1(trace, selected_ctas: Sequence[int]) -> list[TraceRange]: - values = trace.detach().cpu().tolist() - ranges: list[TraceRange] = [] - for cta in selected_ctas: - for dispatch_warp in range(2): - fields = values[cta][dispatch_warp][0] - track = f"dispatch{dispatch_warp}" - _append_range(ranges, fields, 0, 1, cta, track, "metadata") - _append_range(ranges, fields, 2, 3, cta, track, "remote get") - for fields in values[cta][2]: - _append_range(ranges, fields, 0, 1, cta, "producer", "arrival wait") - _append_range(ranges, fields, 1, 2, cta, "producer", "TMA") - for fields in values[cta][3]: - _append_range(ranges, fields, 3, 0, cta, "math", "stage wait") - _append_range(ranges, fields, 0, 1, cta, "math", "GEMM") - _append_range(ranges, fields, 1, 2, cta, "math", "epilogue") - return ranges - - -def _collect_l2(trace, selected_ctas: Sequence[int]) -> list[TraceRange]: - values = trace.detach().cpu().tolist() - ranges: list[TraceRange] = [] - for cta in selected_ctas: - for fields in values[cta][0]: - _append_range(ranges, fields, 0, 1, cta, "producer", "TMA") - for fields in values[cta][1]: - _append_range(ranges, fields, 3, 0, cta, "math", "stage wait") - _append_range(ranges, fields, 0, 1, cta, "math", "GEMM") - _append_range(ranges, fields, 1, 2, cta, "math", "scatter") - fields = values[cta][2][0] - _append_range(ranges, fields, 0, 1, cta, "reduce", "reduce") - return ranges - - -def _draw_panel(draw, ranges: Sequence[TraceRange], selected_ctas: Sequence[int], tracks: Sequence[str], box) -> None: - from PIL import ImageFont - - left, top, right, bottom = box - font = ImageFont.load_default() - label_width = 128 - axis_left = left + label_width - axis_right = right - 12 - title_height = 42 - row_height = max(16, (bottom - top - title_height - 24) // max(1, len(selected_ctas) * len(tracks))) - timestamps = [value for item in ranges for value in (item.start_ns, item.end_ns)] - if not timestamps: - draw.text((left + 8, top + 8), "No trace events", fill="#111827", font=font) - return - - start_ns = min(timestamps) - end_ns = max(timestamps) - duration_ns = max(1, end_ns - start_ns) - for tick in range(6): - x = axis_left + (axis_right - axis_left) * tick / 5 - draw.line((x, top + title_height, x, bottom - 16), fill="#e5e7eb", width=1) - label = f"{duration_ns * tick / 5000:.1f} us" - draw.text((x - 14, bottom - 14), label, fill="#475569", font=font) - - rows = [(cta, track) for cta in selected_ctas for track in tracks] - row_y = {(cta, track): top + title_height + index * row_height for index, (cta, track) in enumerate(rows)} - for index, (cta, track) in enumerate(rows): - y = row_y[(cta, track)] - if index % 2 == 0: - draw.rectangle((left, y, right, y + row_height - 1), fill="#f8fafc") - draw.text((left + 4, y + 2), f"CTA {cta:03d} {track}", fill="#1f2937", font=font) - - for item in ranges: - key = (item.cta, item.track) - if key not in row_y: - continue - x0 = axis_left + (item.start_ns - start_ns) * (axis_right - axis_left) / duration_ns - x1 = axis_left + (item.end_ns - start_ns) * (axis_right - axis_left) / duration_ns - y0 = row_y[key] + 2 - y1 = y0 + max(5, row_height - 5) - draw.rectangle((x0, y0, max(x0 + 1, x1), y1), fill=PHASE_COLORS[item.phase]) - - -def save_pipeline_trace_png( - l1_trace, - l2_trace, - output_path: str | Path, - selected_ctas: Iterable[int], - title: str, -) -> Path: - from PIL import Image, ImageDraw, ImageFont - - selected = tuple(dict.fromkeys(int(cta) for cta in selected_ctas)) - l1_ranges = _collect_l1(l1_trace, selected) - l2_ranges = _collect_l2(l2_trace, selected) - width = 1900 - rows = max(len(selected) * 4, len(selected) * 3) - height = max(420, 112 + rows * 22) - image = Image.new("RGB", (width, height), "white") - draw = ImageDraw.Draw(image) - font = ImageFont.load_default() - draw.text((18, 12), title, fill="#111827", font=font) - - legend_x = 18 - for phase, color in PHASE_COLORS.items(): - draw.rectangle((legend_x, 34, legend_x + 12, 46), fill=color) - draw.text((legend_x + 17, 34), phase, fill="#334155", font=font) - legend_x += 17 + 7 * len(phase) + 18 - - panel_top = 62 - panel_bottom = height - 8 - panel_mid = width // 2 - draw.text((18, panel_top), "L1: dispatch / TMA / GEMM / SwiGLU", fill="#111827", font=font) - draw.text((panel_mid + 8, panel_top), "L2: TMA / GEMM / scatter / reduce", fill="#111827", font=font) - _draw_panel( - draw, - l1_ranges, - selected, - ("dispatch0", "dispatch1", "producer", "math"), - (12, panel_top + 18, panel_mid - 4, panel_bottom), - ) - _draw_panel(draw, l2_ranges, selected, ("producer", "math", "reduce"), (panel_mid + 4, panel_top + 18, width - 12, panel_bottom)) - - output = Path(output_path).expanduser().absolute() - output.parent.mkdir(parents=True, exist_ok=True) - image.save(output) - return output From 917cc958ac5989071ee2db7e92906ccfb4fbd60e Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Wed, 19 Aug 2026 11:22:44 +0800 Subject: [PATCH 20/30] Optimize SM90 MegaMoE two-kernel path --- .../mega_moe/example_sm90_fp8_mega_moe.py | 749 ++++-------------- 1 file changed, 164 insertions(+), 585 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 4bc28a9864..eb11530027 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -83,38 +83,22 @@ def select_manual_warp_configs( expected_tokens = (routed_tokens + num_experts_per_rank - 1) // num_experts_per_rank num_m_blocks = (expected_tokens + 63) // 64 blocks_per_expert = num_m_blocks * (2 * intermediate_hidden // 256) - requested = min( - num_experts_per_rank, - (2 * num_sms + blocks_per_expert - 1) // blocks_per_expert, - ) + requested = min(num_experts_per_rank, (2 * num_sms + blocks_per_expert - 1) // blocks_per_expert) if blocks_per_expert < num_sms: max_candidate = min(num_experts_per_rank, 2 * requested) requested = max( range(requested, max_candidate + 1), - key=lambda candidate: ( - 1.0 - if num_experts_per_rank % candidate == 0 - else (num_experts_per_rank % candidate) / candidate - ), + key=lambda candidate: 1.0 if num_experts_per_rank % candidate == 0 else (num_experts_per_rank % candidate) / candidate, ) - generic_experts_per_wave = normalize_experts_per_wave( - num_experts_per_rank, requested - ) + generic_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, requested) l1_experts_per_wave = l2_experts_per_wave = generic_experts_per_wave if high_sm and shape_family == "compact": if routed_tokens <= 32 * num_experts_per_rank: l1_stages = l2_stages = 3 - l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave( - num_experts_per_rank, 4 - ) - elif ( - 128 * num_experts_per_rank < routed_tokens <= 256 * num_experts_per_rank - or routed_tokens > 1024 * num_experts_per_rank - ): + l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 4) + elif 128 * num_experts_per_rank < routed_tokens <= 256 * num_experts_per_rank or routed_tokens > 1024 * num_experts_per_rank: l1_stages = l2_stages = 4 - l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave( - num_experts_per_rank, 32 - ) + l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 32) elif high_sm and shape_family == "wide": # BN512/BK256 are profitable in the CUDA kernel, but the manually # tuned TileScale BN256/BK128 path is faster for the current WGMMA @@ -123,30 +107,18 @@ def select_manual_warp_configs( if routed_tokens <= 24 * num_experts_per_rank: # CUDA selects 16 experts here, while TileScale's direct TIR # scheduler is faster with a shorter four-expert scan on H200. - l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave( - num_experts_per_rank, 4 - ) + l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 4) elif 24 * num_experts_per_rank < routed_tokens <= 48 * num_experts_per_rank: l1_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 8) l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 48) elif routed_tokens > 48 * num_experts_per_rank: - l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave( - num_experts_per_rank, 16 - ) + l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 16) common = {"block_m": 64, "block_n": 256, "block_k": 128, "threads": 384} return ( shape_family, - { - **common, - "pipeline_stages": l1_stages, - "num_experts_per_wave": l1_experts_per_wave, - }, - { - **common, - "pipeline_stages": l2_stages, - "num_experts_per_wave": l2_experts_per_wave, - }, + {**common, "pipeline_stages": l1_stages, "num_experts_per_wave": l1_experts_per_wave}, + {**common, "pipeline_stages": l2_stages, "num_experts_per_wave": l2_experts_per_wave}, ) @@ -161,13 +133,7 @@ def per_token_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: def block_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: groups, n, k = x.shape - x_view = x.float().view( - groups, - n // SCALE_GRANULARITY, - SCALE_GRANULARITY, - k // SCALE_GRANULARITY, - SCALE_GRANULARITY, - ) + x_view = x.float().view(groups, n // SCALE_GRANULARITY, SCALE_GRANULARITY, k // SCALE_GRANULARITY, SCALE_GRANULARITY) amax = x_view.abs().amax(dim=(-1, -3)).clamp(1e-4) scale = amax / FP8_MAX x_fp8 = (x_view / scale.unsqueeze(-1).unsqueeze(-3)).to(torch.float8_e4m3fn) @@ -189,17 +155,9 @@ def dequantize_per_token(x: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: def dequantize_block(x: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: groups, n, k = x.shape - x_view = x.float().view( - groups, - n // SCALE_GRANULARITY, - SCALE_GRANULARITY, - k // SCALE_GRANULARITY, - SCALE_GRANULARITY, - ) + x_view = x.float().view(groups, n // SCALE_GRANULARITY, SCALE_GRANULARITY, k // SCALE_GRANULARITY, SCALE_GRANULARITY) return (x_view * scale.unsqueeze(-1).unsqueeze(-3)).view(groups, n, k) - - def fused_l1_swiglu_manual_warp_kernel( num_tokens: int, hidden: int, @@ -228,7 +186,8 @@ def fused_l1_swiglu_manual_warp_kernel( max_tiles_per_expert_wave = num_experts_per_wave * num_m_blocks * num_n_blocks max_tile_rounds_per_expert_wave = T.ceildiv(max_tiles_per_expert_wave, num_sms) num_k_blocks = hidden // block_k - # CTA roles: warps 0-1 dispatch, warps 2-3 TMA, then WGMMA warpgroups. + # WG0 is the frontend: warps 0-1 dispatch routes and warps 2-3 issue TMA. + # WG1+ are WGMMA consumers; each owns an N fragment of the CTA tile. warp_size = 32 warpgroup_size = 128 dispatch_warps = 2 @@ -262,9 +221,7 @@ def main( route_counts: T.Tensor((num_ranks, num_experts), T.int32), recv_counts: T.Tensor((num_experts_per_rank,), T.int32), route_slots: T.Tensor((num_tokens, num_topk), T.int32), - arrivals: T.Tensor( - (num_experts_per_rank, T.ceildiv(capacity, block_m)), T.uint32 - ), + arrivals: T.Tensor((num_experts_per_rank, T.ceildiv(capacity, block_m)), T.uint32), recv_x: T.Tensor((num_experts_per_rank, capacity, hidden), T.float8_e4m3fn), recv_x_sf: T.Tensor((num_experts_per_rank, capacity, num_scale_groups), T.float32), recv_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), @@ -272,15 +229,9 @@ def main( src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), src_topk: T.Tensor((num_experts_per_rank, capacity), T.int32), l1_weight: T.Tensor((num_experts_per_rank, l1_n, hidden), T.float8_e4m3fn), - l1_weight_sf: T.Tensor( - (num_experts_per_rank, l1_n // SCALE_GRANULARITY, hidden // SCALE_GRANULARITY), - T.float32, - ), + l1_weight_sf: T.Tensor((num_experts_per_rank, l1_n // SCALE_GRANULARITY, hidden // SCALE_GRANULARITY), T.float32), l2_x: T.Tensor((num_experts_per_rank, capacity, l1_n // 2), T.float8_e4m3fn), - l2_x_sf: T.Tensor( - (num_experts_per_rank, capacity, l1_n // (2 * SCALE_GRANULARITY)), - T.float32, - ), + l2_x_sf: T.Tensor((num_experts_per_rank, capacity, l1_n // (2 * SCALE_GRANULARITY)), T.float32), barrier: T.Tensor((num_ranks,), T.int32), ): T.annotate_pass_configs({tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) @@ -289,18 +240,10 @@ def main( src_rank = T.alloc_local((1,), T.int32) src_rank[0] = T.get_rank() - a_shared = T.alloc_shared( - (pipeline_stages, block_m, block_k), - T.float8_e4m3fn, - ) - b_shared = T.alloc_shared( - (pipeline_stages, block_n, block_k), - T.float8_e4m3fn, - ) + a_shared = T.alloc_shared((pipeline_stages, block_m, block_k), T.float8_e4m3fn) + b_shared = T.alloc_shared((pipeline_stages, block_n, block_k), T.float8_e4m3fn) out_shared = T.alloc_shared((block_m, block_n // 2), T.float8_e4m3fn) - stage_barriers = T.alloc_barrier( - [producer_threads] * pipeline_stages + [num_math_threads] * pipeline_stages - ) + stage_barriers = T.alloc_barrier([producer_threads] * pipeline_stages + [num_math_threads] * pipeline_stages) if bid == 0: if tx < route_threads: @@ -313,34 +256,20 @@ def main( assign_topk = assign_route % num_topk assign_expert = topk_idx[assign_token, assign_topk] if assign_expert >= 0 and assign_expert < num_experts: - route_slots[assign_token, assign_topk] = T.atomic_add( - route_counts[src_rank[0], assign_expert], - 1, - memory_order="relaxed", - return_prev=True, - ) + route_slots[assign_token, assign_topk] = T.atomic_add(route_counts[src_rank[0], assign_expert], 1, memory_order="relaxed", return_prev=True) else: route_slots[assign_token, assign_topk] = -1 T.sync_threads(7, route_threads) - for publish_idx in T.serial( - tx, num_experts * num_ranks, route_threads - ): - publish_rank = publish_idx // num_experts - publish_expert = publish_idx % num_experts + publish_warp = tx // warp_size + for publish_rank in T.serial(publish_warp, num_ranks, route_threads // warp_size): if publish_rank != src_rank[0]: - T.st( - route_counts[src_rank[0], publish_expert], - route_counts[src_rank[0], publish_expert], - dst_pe=publish_rank, - ) + T.put_warp(T.address_of(route_counts[src_rank[0], 0]), T.address_of(route_counts[src_rank[0], 0]), num_experts, dst_pe=publish_rank, unroll_factor=8) T.barrier_blocks(barrier[0]) if tx < route_threads: - for count_local_expert in T.serial( - tx, num_experts_per_rank, route_threads - ): + for count_local_expert in T.serial(tx, num_experts_per_rank, route_threads): recv_count = T.alloc_var(T.int32, init=0) recv_expert = src_rank[0] * num_experts_per_rank + count_local_expert for count_rank in T.serial(num_ranks): @@ -351,10 +280,7 @@ def main( prefix_token = prefix_route // num_topk prefix_topk = prefix_route % num_topk prefix_expert = topk_idx[prefix_token, prefix_topk] - prefix_slot = T.alloc_var( - T.int32, - init=route_slots[prefix_token, prefix_topk], - ) + prefix_slot = T.alloc_var(T.int32, init=route_slots[prefix_token, prefix_topk]) if prefix_token < num_tokens and prefix_expert >= 0 and prefix_expert < num_experts and prefix_slot >= 0: for prefix_rank in T.serial(num_ranks): if prefix_rank < src_rank[0]: @@ -371,11 +297,8 @@ def main( dispatch_warp = tx // warp_size dispatch_lane = tx % warp_size if tx < dispatch_threads: - for metadata_route in T.serial( - bid * dispatch_threads + tx, - num_routes, - num_sms * dispatch_threads, - ): + # WG0 warps 0-1 publish route metadata and pull remote activation rows. + for metadata_route in T.serial(bid * dispatch_threads + tx, num_routes, num_sms * dispatch_threads): metadata_token = metadata_route // num_topk metadata_topk = metadata_route % num_topk metadata_expert = topk_idx[metadata_token, metadata_topk] @@ -383,112 +306,48 @@ def main( if metadata_expert >= 0 and metadata_slot >= 0 and metadata_slot < capacity: metadata_rank = metadata_expert // num_experts_per_rank metadata_local_expert = metadata_expert % num_experts_per_rank - T.st( - recv_weights[metadata_local_expert, metadata_slot], - topk_weights[metadata_token, metadata_topk], - dst_pe=metadata_rank, - ) - T.st( - src_tokens[metadata_local_expert, metadata_slot], - metadata_token, - dst_pe=metadata_rank, - ) - T.st( - src_topk[metadata_local_expert, metadata_slot], - metadata_topk, - dst_pe=metadata_rank, - ) - T.st( - src_ranks[metadata_local_expert, metadata_slot], - src_rank[0], - scope="sys", - sem="release", - dst_pe=metadata_rank, - ) - - for pull_idx in T.serial( - bid * dispatch_warps + dispatch_warp, - num_experts_per_rank * capacity, - num_sms * dispatch_warps, - ): + T.st(recv_weights[metadata_local_expert, metadata_slot], topk_weights[metadata_token, metadata_topk], dst_pe=metadata_rank) + T.st(src_tokens[metadata_local_expert, metadata_slot], metadata_token, dst_pe=metadata_rank) + T.st(src_topk[metadata_local_expert, metadata_slot], metadata_topk, dst_pe=metadata_rank) + T.st(src_ranks[metadata_local_expert, metadata_slot], src_rank[0], scope="sys", sem="release", dst_pe=metadata_rank) + + for pull_idx in T.serial(bid * dispatch_warps + dispatch_warp, num_experts_per_rank * capacity, num_sms * dispatch_warps): pull_expert = pull_idx // capacity pull_slot = pull_idx % capacity if pull_slot < recv_counts[pull_expert]: if dispatch_lane == dispatch_leader_lane: - T.wait_ge( - src_ranks[pull_expert, pull_slot], - 0, - scope=T.WaitScope.SYS, - semantics=T.WaitSemantics.ACQUIRE, - ) + T.wait_ge(src_ranks[pull_expert, pull_slot], 0, scope=T.WaitScope.SYS, semantics=T.WaitSemantics.ACQUIRE) T.sync_warp() pull_rank = src_ranks[pull_expert, pull_slot] pull_token = src_tokens[pull_expert, pull_slot] - T.get_warp( - T.address_of(x[pull_token, 0]), - T.address_of(recv_x[pull_expert, pull_slot, 0]), - hidden, - src_pe=pull_rank, - unroll_factor=8, - ) - T.get_warp( - T.address_of(x_sf[pull_token, 0]), - T.address_of(recv_x_sf[pull_expert, pull_slot, 0]), - num_scale_groups, - src_pe=pull_rank, - unroll_factor=8, - ) + T.get_warp(T.address_of(x[pull_token, 0]), T.address_of(recv_x[pull_expert, pull_slot, 0]), hidden, src_pe=pull_rank, unroll_factor=8) + T.get_warp(T.address_of(x_sf[pull_token, 0]), T.address_of(recv_x_sf[pull_expert, pull_slot, 0]), num_scale_groups, src_pe=pull_rank, unroll_factor=8) T.sync_warp() if dispatch_lane == dispatch_leader_lane: - T.atom_add( - arrivals[pull_expert, pull_slot // block_m], - 1, - scope="gpu", - sem="release", - ) + T.atom_add(arrivals[pull_expert, pull_slot // block_m], 1, scope="gpu", sem="release") if tx >= producer_begin and tx < producer_end: + # WG0 warps 2-3 keep the shared-memory pipeline filled with TMA. producer_step = T.alloc_var(T.int32, init=0) producer_wave_tile = T.alloc_var(T.int32, init=bid) - for producer_schedule_step in T.serial( - num_expert_waves * max_tile_rounds_per_expert_wave - ): - producer_expert_wave = ( - producer_schedule_step // max_tile_rounds_per_expert_wave - ) - producer_tile_round = ( - producer_schedule_step % max_tile_rounds_per_expert_wave - ) + for producer_schedule_step in T.serial(num_expert_waves * max_tile_rounds_per_expert_wave): + producer_expert_wave = producer_schedule_step // max_tile_rounds_per_expert_wave + producer_tile_round = producer_schedule_step % max_tile_rounds_per_expert_wave producer_wave_begin = producer_expert_wave * num_experts_per_wave producer_wave_num_tiles = T.alloc_var(T.int32, init=0) for producer_wave_expert_offset in T.serial(num_experts_per_wave): - producer_wave_expert = ( - producer_wave_begin + producer_wave_expert_offset - ) - producer_wave_num_tiles += ( - T.ceildiv(T.min(recv_counts[producer_wave_expert], capacity), block_m) - * num_n_blocks - ) + producer_wave_expert = producer_wave_begin + producer_wave_expert_offset + producer_wave_num_tiles += T.ceildiv(T.min(recv_counts[producer_wave_expert], capacity), block_m) * num_n_blocks if producer_wave_tile < producer_wave_num_tiles: - producer_tile_offset = T.alloc_var( - T.int32, init=producer_wave_tile - ) + producer_tile_offset = T.alloc_var(T.int32, init=producer_wave_tile) producer_expert = T.alloc_var(T.int32, init=-1) producer_m = T.alloc_var(T.int32, init=0) producer_n = T.alloc_var(T.int32, init=0) - for producer_wave_expert_offset in T.serial( - num_experts_per_wave - ): - producer_wave_expert = ( - producer_wave_begin + producer_wave_expert_offset - ) - producer_expert_m_blocks = T.ceildiv( - T.min(recv_counts[producer_wave_expert], capacity), block_m - ) - producer_expert_tiles = ( - producer_expert_m_blocks * num_n_blocks - ) + for producer_wave_expert_offset in T.serial(num_experts_per_wave): + producer_wave_expert = producer_wave_begin + producer_wave_expert_offset + producer_expert_m_blocks = T.ceildiv(T.min(recv_counts[producer_wave_expert], capacity), block_m) + producer_expert_tiles = producer_expert_m_blocks * num_n_blocks if producer_expert < 0: if producer_tile_offset < producer_expert_tiles: producer_expert = producer_wave_expert @@ -498,25 +357,14 @@ def main( producer_tile_offset -= producer_expert_tiles if producer_expert >= 0 and producer_n * block_n < l1_n: - producer_arrivals = T.min( - block_m, - recv_counts[producer_expert] - producer_m * block_m, - ) + producer_arrivals = T.min(block_m, recv_counts[producer_expert] - producer_m * block_m) if tx == producer_begin: - T.wait_ge( - arrivals[producer_expert, producer_m], - producer_arrivals, - scope=T.WaitScope.GPU, - semantics=T.WaitSemantics.ACQUIRE, - ) + T.wait_ge(arrivals[producer_expert, producer_m], producer_arrivals, scope=T.WaitScope.GPU, semantics=T.WaitSemantics.ACQUIRE) T.sync_threads(5, producer_threads) for producer_k in T.serial(num_k_blocks): producer_stage = (producer_step + producer_k) % pipeline_stages producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 - T.mbarrier_wait_parity( - stage_barriers[pipeline_stages + producer_stage], - producer_phase ^ 1, - ) + T.mbarrier_wait_parity(stage_barriers[pipeline_stages + producer_stage], producer_phase ^ 1) T.tma_copy( recv_x[ producer_expert, @@ -547,20 +395,15 @@ def main( producer_step += num_k_blocks producer_wave_tile += num_sms - if ( - producer_tile_round == max_tile_rounds_per_expert_wave - 1 - and producer_wave_num_tiles > 0 - ): + if producer_tile_round == max_tile_rounds_per_expert_wave - 1 and producer_wave_num_tiles > 0: producer_wave_tile -= producer_wave_num_tiles elif tx >= math_begin: + # WG1+ consume TMA stages, run WGMMA, and emit quantized L1 rows. partial = T.alloc_fragment((block_m, block_n), T.float32) accum = T.alloc_fragment((block_m, block_n), T.bfloat16) gate = T.alloc_fragment((block_m, block_n // 2), T.float32) - gate_grouped = T.reshape( - gate, - (block_m, num_output_scale_groups, SCALE_GRANULARITY), - ) + gate_grouped = T.reshape(gate, (block_m, num_output_scale_groups, SCALE_GRANULARITY)) up = T.alloc_fragment((block_m, block_n // 2), T.float32) amax = T.alloc_fragment((block_m, num_output_scale_groups), T.float32) scale = T.alloc_fragment((block_m, num_output_scale_groups), T.float32) @@ -570,45 +413,24 @@ def main( consumer_step = T.alloc_var(T.int32, init=0) consumer_wave_tile = T.alloc_var(T.int32, init=bid) - for consumer_schedule_step in T.serial( - num_expert_waves * max_tile_rounds_per_expert_wave - ): - consumer_expert_wave = ( - consumer_schedule_step // max_tile_rounds_per_expert_wave - ) - consumer_tile_round = ( - consumer_schedule_step % max_tile_rounds_per_expert_wave - ) + for consumer_schedule_step in T.serial(num_expert_waves * max_tile_rounds_per_expert_wave): + consumer_expert_wave = consumer_schedule_step // max_tile_rounds_per_expert_wave + consumer_tile_round = consumer_schedule_step % max_tile_rounds_per_expert_wave consumer_wave_begin = consumer_expert_wave * num_experts_per_wave consumer_wave_num_tiles = T.alloc_var(T.int32, init=0) for consumer_wave_expert_offset in T.serial(num_experts_per_wave): - consumer_wave_expert = ( - consumer_wave_begin + consumer_wave_expert_offset - ) - consumer_wave_num_tiles += ( - T.ceildiv(T.min(recv_counts[consumer_wave_expert], capacity), block_m) - * num_n_blocks - ) + consumer_wave_expert = consumer_wave_begin + consumer_wave_expert_offset + consumer_wave_num_tiles += T.ceildiv(T.min(recv_counts[consumer_wave_expert], capacity), block_m) * num_n_blocks if consumer_wave_tile < consumer_wave_num_tiles: - consumer_tile_offset = T.alloc_var( - T.int32, init=consumer_wave_tile - ) + consumer_tile_offset = T.alloc_var(T.int32, init=consumer_wave_tile) consumer_expert = T.alloc_var(T.int32, init=-1) consumer_m = T.alloc_var(T.int32, init=0) consumer_n = T.alloc_var(T.int32, init=0) - for consumer_wave_expert_offset in T.serial( - num_experts_per_wave - ): - consumer_wave_expert = ( - consumer_wave_begin + consumer_wave_expert_offset - ) - consumer_expert_m_blocks = T.ceildiv( - T.min(recv_counts[consumer_wave_expert], capacity), block_m - ) - consumer_expert_tiles = ( - consumer_expert_m_blocks * num_n_blocks - ) + for consumer_wave_expert_offset in T.serial(num_experts_per_wave): + consumer_wave_expert = consumer_wave_begin + consumer_wave_expert_offset + consumer_expert_m_blocks = T.ceildiv(T.min(recv_counts[consumer_wave_expert], capacity), block_m) + consumer_expert_tiles = consumer_expert_m_blocks * num_n_blocks if consumer_expert < 0: if consumer_tile_offset < consumer_expert_tiles: consumer_expert = consumer_wave_expert @@ -623,16 +445,8 @@ def main( for consumer_k in T.serial(num_k_blocks): consumer_stage = (consumer_step + consumer_k) % pipeline_stages consumer_phase = ((consumer_step + consumer_k) // pipeline_stages) & 1 - T.mbarrier_wait_parity( - stage_barriers[consumer_stage], - consumer_phase, - ) - T.gemm( - a_shared[consumer_stage, :, :], - b_shared[consumer_stage, :, :], - partial, - transpose_B=True, - ) + T.mbarrier_wait_parity(stage_barriers[consumer_stage], consumer_phase) + T.gemm(a_shared[consumer_stage, :, :], b_shared[consumer_stage, :, :], partial, transpose_B=True) for i in T.Parallel(block_m): act_scale[i] = recv_x_sf[ consumer_expert, @@ -653,18 +467,9 @@ def main( consumer_k, ] for i, j in T.Parallel(block_m, block_n): - accum[i, j] = ( - T.cast(partial[i, j], T.bfloat16) - * T.cast( - act_scale[i] - * weight_scale[ - 2 * (j // (2 * SCALE_GRANULARITY)) - + (j % 16) // 8 - ], - T.bfloat16, - ) - + accum[i, j] - ) + accum[i, j] = T.cast(partial[i, j], T.bfloat16) * T.cast( + act_scale[i] * weight_scale[2 * (j // (2 * SCALE_GRANULARITY)) + (j % 16) // 8], T.bfloat16 + ) + accum[i, j] T.clear(partial) T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) consumer_step += num_k_blocks @@ -686,10 +491,7 @@ def main( ] ) T.reduce_absmax(gate_grouped, amax, dim=2) - for i, scale_group in T.Parallel( - block_m, - num_output_scale_groups, - ): + for i, scale_group in T.Parallel(block_m, num_output_scale_groups): scale[i, scale_group] = T.max(amax[i, scale_group], 1e-4) / FP8_MAX l2_x_sf[ consumer_expert, @@ -697,11 +499,7 @@ def main( consumer_n * num_output_scale_groups + scale_group, ] = scale[i, scale_group] for i, j in T.Parallel(block_m, block_n // 2): - gate[i, j] = T.clamp( - gate[i, j] / scale[i, j // SCALE_GRANULARITY], - -FP8_MAX, - FP8_MAX, - ) + gate[i, j] = T.clamp(gate[i, j] / scale[i, j // SCALE_GRANULARITY], -FP8_MAX, FP8_MAX) T.copy(gate, quant_fp8) T.copy(quant_fp8, out_shared) T.copy( @@ -714,10 +512,7 @@ def main( ) consumer_wave_tile += num_sms - if ( - consumer_tile_round == max_tile_rounds_per_expert_wave - 1 - and consumer_wave_num_tiles > 0 - ): + if consumer_tile_round == max_tile_rounds_per_expert_wave - 1 and consumer_wave_num_tiles > 0: consumer_wave_tile -= consumer_wave_num_tiles return main @@ -753,8 +548,9 @@ def fused_l2_scatter_reduce_manual_warp_kernel( num_reduce_n_blocks = T.ceildiv(hidden, reduce_block_h) num_reduce_m_blocks = T.ceildiv(num_tokens, reduce_block_m) num_reduce_tiles = num_reduce_n_blocks * num_reduce_m_blocks - # GEMM roles: warps 0-1 idle, warps 2-3 TMA, then two WGMMA - # warpgroups. After the grid barrier, warps 0-3 perform top-k reduction. + # WG0 is phase-specialized: warps 0-1 are idle during GEMM, while warps + # 2-3 issue TMA. WG1-2 run WGMMA and direct scatter; after the grid + # barrier, all WG0 warps reduce the top-k slots. warp_size = 32 warpgroup_size = 128 reduce_warps = 4 @@ -768,38 +564,17 @@ def fused_l2_scatter_reduce_manual_warp_kernel( producer_end = producer_begin + producer_threads math_begin = math_begin_warp * warp_size num_math_threads = math_warpgroups * warpgroup_size - math_warps = num_math_threads // warp_size - rows_per_math_warp = block_m // math_warps assert producer_end == math_begin == reduce_threads - assert block_m % math_warps == 0 + assert block_m == 64 + assert block_n % (math_warpgroups * 8) == 0 assert threads == math_begin + num_math_threads @T.prim_func def main( - a: T.Tensor( - (num_experts_per_rank, capacity, intermediate_hidden), - T.float8_e4m3fn, - ), - b: T.Tensor( - (num_experts_per_rank, hidden, intermediate_hidden), - T.float8_e4m3fn, - ), - a_sf: T.Tensor( - ( - num_experts_per_rank, - capacity, - intermediate_hidden // SCALE_GRANULARITY, - ), - T.float32, - ), - b_sf: T.Tensor( - ( - num_experts_per_rank, - hidden // SCALE_GRANULARITY, - intermediate_hidden // SCALE_GRANULARITY, - ), - T.float32, - ), + a: T.Tensor((num_experts_per_rank, capacity, intermediate_hidden), T.float8_e4m3fn), + b: T.Tensor((num_experts_per_rank, hidden, intermediate_hidden), T.float8_e4m3fn), + a_sf: T.Tensor((num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), T.float32), + b_sf: T.Tensor((num_experts_per_rank, hidden // SCALE_GRANULARITY, intermediate_hidden // SCALE_GRANULARITY), T.float32), recv_counts: T.Tensor((num_experts_per_rank,), T.int32), src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), @@ -810,20 +585,11 @@ def main( ): with T.Kernel(num_sms, threads=threads) as bid: tx = T.get_thread_binding() - a_shared = T.alloc_shared( - (pipeline_stages, block_m, block_k), - T.float8_e4m3fn, - ) - b_shared = T.alloc_shared( - (pipeline_stages, block_n, block_k), - T.float8_e4m3fn, - ) + a_shared = T.alloc_shared((pipeline_stages, block_m, block_k), T.float8_e4m3fn) + b_shared = T.alloc_shared((pipeline_stages, block_n, block_k), T.float8_e4m3fn) a_sf_shared = T.alloc_shared((pipeline_stages, block_m), T.float32) - out_shared = T.alloc_shared((block_m, block_n), T.bfloat16) reduce_shared = T.alloc_shared((reduce_block_m, reduce_block_h), T.bfloat16) - stage_barriers = T.alloc_barrier( - [producer_threads] * pipeline_stages + [num_math_threads] * pipeline_stages - ) + stage_barriers = T.alloc_barrier([producer_threads] * pipeline_stages + [num_math_threads] * pipeline_stages) if tx < math_begin: T.dec_max_nreg(48) @@ -831,47 +597,27 @@ def main( T.inc_max_nreg(208) if tx >= producer_begin and tx < producer_end: + # WG0 warps 2-3 keep the L2 TMA stages filled. producer_step = T.alloc_var(T.int32, init=0) producer_wave_tile = T.alloc_var(T.int32, init=bid) - for producer_schedule_step in T.serial( - num_expert_waves * max_tile_rounds_per_expert_wave - ): - producer_expert_wave = ( - producer_schedule_step // max_tile_rounds_per_expert_wave - ) - producer_tile_round = ( - producer_schedule_step % max_tile_rounds_per_expert_wave - ) + for producer_schedule_step in T.serial(num_expert_waves * max_tile_rounds_per_expert_wave): + producer_expert_wave = producer_schedule_step // max_tile_rounds_per_expert_wave + producer_tile_round = producer_schedule_step % max_tile_rounds_per_expert_wave producer_wave_begin = producer_expert_wave * num_experts_per_wave producer_wave_num_tiles = T.alloc_var(T.int32, init=0) for producer_wave_expert_offset in T.serial(num_experts_per_wave): - producer_wave_expert = ( - producer_wave_begin + producer_wave_expert_offset - ) - producer_wave_num_tiles += ( - T.ceildiv(T.min(recv_counts[producer_wave_expert], capacity), block_m) - * num_n_blocks - ) + producer_wave_expert = producer_wave_begin + producer_wave_expert_offset + producer_wave_num_tiles += T.ceildiv(T.min(recv_counts[producer_wave_expert], capacity), block_m) * num_n_blocks if producer_wave_tile < producer_wave_num_tiles: - producer_tile_offset = T.alloc_var( - T.int32, init=producer_wave_tile - ) + producer_tile_offset = T.alloc_var(T.int32, init=producer_wave_tile) producer_expert = T.alloc_var(T.int32, init=-1) producer_m = T.alloc_var(T.int32, init=0) producer_n = T.alloc_var(T.int32, init=0) - for producer_wave_expert_offset in T.serial( - num_experts_per_wave - ): - producer_wave_expert = ( - producer_wave_begin + producer_wave_expert_offset - ) - producer_expert_m_blocks = T.ceildiv( - T.min(recv_counts[producer_wave_expert], capacity), block_m - ) - producer_expert_tiles = ( - producer_expert_m_blocks * num_n_blocks - ) + for producer_wave_expert_offset in T.serial(num_experts_per_wave): + producer_wave_expert = producer_wave_begin + producer_wave_expert_offset + producer_expert_m_blocks = T.ceildiv(T.min(recv_counts[producer_wave_expert], capacity), block_m) + producer_expert_tiles = producer_expert_m_blocks * num_n_blocks if producer_expert < 0: if producer_tile_offset < producer_expert_tiles: producer_expert = producer_wave_expert @@ -889,10 +635,7 @@ def main( for producer_k in T.serial(num_k_blocks): producer_stage = (producer_step + producer_k) % pipeline_stages producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 - T.mbarrier_wait_parity( - stage_barriers[pipeline_stages + producer_stage], - producer_phase ^ 1, - ) + T.mbarrier_wait_parity(stage_barriers[pipeline_stages + producer_stage], producer_phase ^ 1) T.tma_copy( a[ producer_expert, @@ -911,22 +654,16 @@ def main( b_shared[producer_stage, :, :], barrier=stage_barriers[producer_stage], ) - a_sf_shared[producer_stage, tx - producer_begin] = a_sf[ - producer_expert, - producer_m * block_m + tx - producer_begin, - producer_k, - ] + a_sf_shared[producer_stage, tx - producer_begin] = a_sf[producer_expert, producer_m * block_m + tx - producer_begin, producer_k] T.mbarrier_arrive(stage_barriers[producer_stage]) producer_step += num_k_blocks producer_wave_tile += num_sms - if ( - producer_tile_round == max_tile_rounds_per_expert_wave - 1 - and producer_wave_num_tiles > 0 - ): + if producer_tile_round == max_tile_rounds_per_expert_wave - 1 and producer_wave_num_tiles > 0: producer_wave_tile -= producer_wave_num_tiles elif tx >= math_begin: + # WG1-2 run WGMMA and scatter their BF16 column pairs remotely. partial = T.alloc_fragment((block_m, block_n), T.float32) accum = T.alloc_fragment((block_m, block_n), T.bfloat16) act_scale = T.alloc_fragment((block_m,), T.float32) @@ -937,45 +674,24 @@ def main( scatter_dst_token = T.alloc_var(T.int32, init=0) scatter_dst_topk = T.alloc_var(T.int32, init=0) - for consumer_schedule_step in T.serial( - num_expert_waves * max_tile_rounds_per_expert_wave - ): - consumer_expert_wave = ( - consumer_schedule_step // max_tile_rounds_per_expert_wave - ) - consumer_tile_round = ( - consumer_schedule_step % max_tile_rounds_per_expert_wave - ) + for consumer_schedule_step in T.serial(num_expert_waves * max_tile_rounds_per_expert_wave): + consumer_expert_wave = consumer_schedule_step // max_tile_rounds_per_expert_wave + consumer_tile_round = consumer_schedule_step % max_tile_rounds_per_expert_wave consumer_wave_begin = consumer_expert_wave * num_experts_per_wave consumer_wave_num_tiles = T.alloc_var(T.int32, init=0) for consumer_wave_expert_offset in T.serial(num_experts_per_wave): - consumer_wave_expert = ( - consumer_wave_begin + consumer_wave_expert_offset - ) - consumer_wave_num_tiles += ( - T.ceildiv(T.min(recv_counts[consumer_wave_expert], capacity), block_m) - * num_n_blocks - ) + consumer_wave_expert = consumer_wave_begin + consumer_wave_expert_offset + consumer_wave_num_tiles += T.ceildiv(T.min(recv_counts[consumer_wave_expert], capacity), block_m) * num_n_blocks if consumer_wave_tile < consumer_wave_num_tiles: - consumer_tile_offset = T.alloc_var( - T.int32, init=consumer_wave_tile - ) + consumer_tile_offset = T.alloc_var(T.int32, init=consumer_wave_tile) consumer_expert = T.alloc_var(T.int32, init=-1) consumer_m = T.alloc_var(T.int32, init=0) consumer_n = T.alloc_var(T.int32, init=0) - for consumer_wave_expert_offset in T.serial( - num_experts_per_wave - ): - consumer_wave_expert = ( - consumer_wave_begin + consumer_wave_expert_offset - ) - consumer_expert_m_blocks = T.ceildiv( - T.min(recv_counts[consumer_wave_expert], capacity), block_m - ) - consumer_expert_tiles = ( - consumer_expert_m_blocks * num_n_blocks - ) + for consumer_wave_expert_offset in T.serial(num_experts_per_wave): + consumer_wave_expert = consumer_wave_begin + consumer_wave_expert_offset + consumer_expert_m_blocks = T.ceildiv(T.min(recv_counts[consumer_wave_expert], capacity), block_m) + consumer_expert_tiles = consumer_expert_m_blocks * num_n_blocks if consumer_expert < 0: if consumer_tile_offset < consumer_expert_tiles: consumer_expert = consumer_wave_expert @@ -995,80 +711,49 @@ def main( for consumer_k in T.serial(num_k_blocks): consumer_stage = (consumer_step + consumer_k) % pipeline_stages consumer_phase = ((consumer_step + consumer_k) // pipeline_stages) & 1 - T.mbarrier_wait_parity( - stage_barriers[consumer_stage], - consumer_phase, - ) - T.gemm( - a_shared[consumer_stage, :, :], - b_shared[consumer_stage, :, :], - partial, - transpose_B=True, - ) + T.mbarrier_wait_parity(stage_barriers[consumer_stage], consumer_phase) + T.gemm(a_shared[consumer_stage, :, :], b_shared[consumer_stage, :, :], partial, transpose_B=True) for i in T.Parallel(block_m): act_scale[i] = a_sf_shared[consumer_stage, i] - weight_scale[0] = b_sf[ - consumer_expert, - consumer_n * 2, - consumer_k, - ] - weight_scale[1] = b_sf[ - consumer_expert, - consumer_n * 2 + 1, - consumer_k, - ] + weight_scale[0] = b_sf[consumer_expert, consumer_n * 2, consumer_k] + weight_scale[1] = b_sf[consumer_expert, consumer_n * 2 + 1, consumer_k] for i, j in T.Parallel(block_m, block_n): - accum[i, j] = ( - T.cast(partial[i, j], T.bfloat16) - * T.cast( - act_scale[i] * weight_scale[j // 128], - T.bfloat16, - ) - + accum[i, j] - ) + accum[i, j] = T.cast(partial[i, j], T.bfloat16) * T.cast( + act_scale[i] * weight_scale[j // 128], T.bfloat16 + ) + accum[i, j] T.clear(partial) T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) consumer_step += num_k_blocks - T.copy(accum, out_shared) - scatter_warp = (tx - math_begin) // warp_size - for row_in_warp in T.serial(rows_per_math_warp): - row = scatter_warp * rows_per_math_warp + row_in_warp + # Map the SM90 M64 WGMMA fragment owners directly to + # packed BF16 remote stores, without a shared-memory epilogue. + scatter_math_thread = tx - math_begin + scatter_wg = scatter_math_thread // warpgroup_size + scatter_warp_in_wg = (scatter_math_thread % warpgroup_size) // warp_size + scatter_lane = scatter_math_thread % warp_size + for scatter_row_half in T.serial(2): + row = scatter_warp_in_wg * 16 + scatter_row_half * 8 + scatter_lane // 4 pool_row = consumer_m * block_m + row if pool_row < recv_counts[consumer_expert]: - if tx % warp_size == 0: - scatter_dst_rank = src_ranks[consumer_expert, pool_row] - scatter_dst_token = src_tokens[consumer_expert, pool_row] - scatter_dst_topk = src_topk[consumer_expert, pool_row] - dst_rank = T.shfl_sync(scatter_dst_rank, 0) - dst_token = T.shfl_sync(scatter_dst_token, 0) - dst_topk = T.shfl_sync(scatter_dst_topk, 0) + scatter_dst_rank = src_ranks[consumer_expert, pool_row] + scatter_dst_token = src_tokens[consumer_expert, pool_row] + scatter_dst_topk = src_topk[consumer_expert, pool_row] if ( - dst_rank >= 0 - and dst_rank < num_ranks - and dst_token >= 0 - and dst_token < num_tokens - and dst_topk >= 0 - and dst_topk < num_topk + scatter_dst_rank >= 0 + and scatter_dst_rank < num_ranks + and scatter_dst_token >= 0 + and scatter_dst_token < num_tokens + and scatter_dst_topk >= 0 + and scatter_dst_topk < num_topk ): - T.put_warp( - T.address_of(out_shared[row, 0]), - T.address_of( - combine[ - dst_token, - dst_topk, - consumer_n * block_n, - ] - ), - block_n, - dst_pe=dst_rank, - unroll_factor=1, - ) + for scatter_col_chunk in T.serial(block_n // math_warpgroups // 8): + scatter_col = scatter_wg * (block_n // math_warpgroups) + scatter_col_chunk * 8 + (scatter_lane % 4) * 2 + scatter_value_lo = T.alloc_var(T.uint16, init=T.reinterpret(accum[row, scatter_col], T.uint16)) + scatter_value_hi = T.alloc_var(T.uint16, init=T.reinterpret(accum[row, scatter_col + 1], T.uint16)) + scatter_value = T.alloc_var(T.uint32, init=T.cast(scatter_value_lo, T.uint32) | (T.cast(scatter_value_hi, T.uint32) << 16)) + T.st(combine[scatter_dst_token, scatter_dst_topk, consumer_n * block_n + scatter_col], scatter_value, dst_pe=scatter_dst_rank) consumer_wave_tile += num_sms - if ( - consumer_tile_round == max_tile_rounds_per_expert_wave - 1 - and consumer_wave_num_tiles > 0 - ): + if consumer_tile_round == max_tile_rounds_per_expert_wave - 1 and consumer_wave_num_tiles > 0: consumer_wave_tile -= consumer_wave_num_tiles T.fence_sys() @@ -1078,6 +763,7 @@ def main( T.sync_grid() if tx < reduce_threads: + # After every remote scatter is visible, WG0 reduces top-k into out. reduce_accum = T.alloc_fragment((reduce_block_m, reduce_block_h), T.float32) for reduce_tile in T.serial(bid, num_reduce_tiles, num_sms): reduce_n = reduce_tile % num_reduce_n_blocks @@ -1086,11 +772,7 @@ def main( for topk_slot in T.serial(num_topk): for i, j in T.Parallel(reduce_block_m, reduce_block_h): if reduce_m * reduce_block_m + i < num_tokens: - reduce_accum[i, j] += combine[ - reduce_m * reduce_block_m + i, - topk_slot, - reduce_n * reduce_block_h + j, - ] + reduce_accum[i, j] += combine[reduce_m * reduce_block_m + i, topk_slot, reduce_n * reduce_block_h + j] T.copy(reduce_accum, reduce_shared) T.copy( reduce_shared, @@ -1116,9 +798,7 @@ def _allocator_size_bytes( fp32 = 4 i32 = 4 weight_bytes = num_experts_per_rank * (2 * intermediate_hidden * hidden * fp8 + hidden * intermediate_hidden * fp8) - weight_scale_bytes = ( - num_experts_per_rank * ((2 * intermediate_hidden // 128) * (hidden // 128) + (hidden // 128) * (intermediate_hidden // 128)) * fp32 - ) + weight_scale_bytes = num_experts_per_rank * ((2 * intermediate_hidden // 128) * (hidden // 128) + (hidden // 128) * (intermediate_hidden // 128)) * fp32 pool_bytes = ( num_experts_per_rank * capacity @@ -1213,28 +893,15 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): assert num_experts > 0 and num_experts % num_local_ranks == 0 assert 0 < num_topk <= min(32, num_experts) num_experts_per_rank = num_experts // num_local_ranks - average_recv = ( - num_tokens * num_local_ranks * num_topk + num_experts - 1 - ) // num_experts - capacity = ( - args.capacity - if args.capacity is not None - else (max(average_recv * 2, 64) + 63) // 64 * 64 - ) + average_recv = (num_tokens * num_local_ranks * num_topk + num_experts - 1) // num_experts + capacity = args.capacity if args.capacity is not None else (max(average_recv * 2, 64) + 63) // 64 * 64 assert capacity >= 64 and capacity % 64 == 0 rank, num_ranks, group = init_dist(local_rank, num_local_ranks) assert rank == local_rank and num_ranks == num_local_ranks num_sms = torch.cuda.get_device_properties(local_rank).multi_processor_count allocator = get_allocator( - size=_allocator_size_bytes( - num_tokens, - hidden, - intermediate_hidden, - num_experts_per_rank, - num_topk, - capacity, - ), + size=_allocator_size_bytes(num_tokens, hidden, intermediate_hidden, num_experts_per_rank, num_topk, capacity), device=f"cuda:{local_rank}", is_distributed=True, local_rank=local_rank, @@ -1243,14 +910,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): use_vmm=True, ) - shape_family, l1_config, l2_config = select_manual_warp_configs( - hidden, - intermediate_hidden, - num_tokens, - num_topk, - num_experts_per_rank, - num_sms, - ) + shape_family, l1_config, l2_config = select_manual_warp_configs(hidden, intermediate_hidden, num_tokens, num_topk, num_experts_per_rank, num_sms) for phase, config in (("l1", l1_config), ("l2", l2_config)): requested = getattr(args, f"{phase}_experts_per_wave", None) if requested is not None: @@ -1258,27 +918,11 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): config["num_experts_per_wave"] = requested kernel_specs = [ fused_l1_swiglu_manual_warp_kernel( - num_tokens, - hidden, - 2 * intermediate_hidden, - num_experts, - num_topk, - num_ranks, - capacity, - num_sms, - activation_clamp=activation_clamp, - **l1_config, + num_tokens, hidden, 2 * intermediate_hidden, num_experts, num_topk, num_ranks, capacity, num_sms, + activation_clamp=activation_clamp, **l1_config, ), fused_l2_scatter_reduce_manual_warp_kernel( - num_tokens, - hidden, - intermediate_hidden, - num_experts_per_rank, - num_topk, - num_ranks, - capacity, - num_sms, - **l2_config, + num_tokens, hidden, intermediate_hidden, num_experts_per_rank, num_topk, num_ranks, capacity, num_sms, **l2_config, ), ] kernels = [tilelang.compile(spec, compile_once=True, compile_group=group) for spec in kernel_specs] @@ -1297,22 +941,8 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): topk_weights_src, topk_idx_src = torch.topk(scores, num_topk, dim=-1, sorted=False) topk_idx_src = topk_idx_src.to(torch.int32) - l1_bf16 = ( - torch.randn( - (num_experts_per_rank, 2 * intermediate_hidden, hidden), - dtype=torch.bfloat16, - device="cuda", - ) - * 0.05 - ) - l2_bf16 = ( - torch.randn( - (num_experts_per_rank, hidden, intermediate_hidden), - dtype=torch.bfloat16, - device="cuda", - ) - * 0.05 - ) + l1_bf16 = torch.randn((num_experts_per_rank, 2 * intermediate_hidden, hidden), dtype=torch.bfloat16, device="cuda") * 0.05 + l2_bf16 = torch.randn((num_experts_per_rank, hidden, intermediate_hidden), dtype=torch.bfloat16, device="cuda") * 0.05 l1_fp8_src, l1_sf_src = block_cast_to_fp8(l1_bf16) l2_fp8_src, l2_sf_src = block_cast_to_fp8(l2_bf16) del scores, l1_bf16, l2_bf16 @@ -1334,27 +964,15 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): route_counts = allocator_tensor((num_ranks, num_experts), torch.int32, allocator=allocator) recv_counts = allocator_tensor((num_experts_per_rank,), torch.int32, allocator=allocator) recv_x = allocator_tensor((num_experts_per_rank, capacity, hidden), torch.float8_e4m3fn, allocator=allocator) - recv_x_sf = allocator_tensor( - (num_experts_per_rank, capacity, hidden // SCALE_GRANULARITY), - torch.float32, - allocator=allocator, - ) + recv_x_sf = allocator_tensor((num_experts_per_rank, capacity, hidden // SCALE_GRANULARITY), torch.float32, allocator=allocator) recv_weights = allocator_tensor((num_experts_per_rank, capacity), torch.float32, allocator=allocator) src_ranks = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) src_tokens = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) src_topk = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) route_slots = allocator_tensor((num_tokens, num_topk), torch.int32, allocator=allocator) - arrivals = allocator_tensor( - (num_experts_per_rank, (capacity + 63) // 64), - torch.uint32, - allocator=allocator, - ) + arrivals = allocator_tensor((num_experts_per_rank, (capacity + 63) // 64), torch.uint32, allocator=allocator) l2_x = allocator_tensor((num_experts_per_rank, capacity, intermediate_hidden), torch.float8_e4m3fn, allocator=allocator) - l2_x_sf = allocator_tensor( - (num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), - torch.float32, - allocator=allocator, - ) + l2_x_sf = allocator_tensor((num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), torch.float32, allocator=allocator) combine = allocator_tensor((num_tokens, num_topk, hidden), torch.bfloat16, allocator=allocator) out = allocator_tensor((num_tokens, hidden), torch.bfloat16, allocator=allocator) @@ -1373,42 +991,17 @@ def reset_state(): def run_pipeline(check_capacity: bool = False): fused_l1( - x, - x_sf, - topk_idx, - topk_weights, - route_counts, - recv_counts, - route_slots, - arrivals, - recv_x, - recv_x_sf, - recv_weights, - src_ranks, - src_tokens, - src_topk, - l1_fp8, - l1_sf, - l2_x, - l2_x_sf, - barrier, + x, x_sf, topk_idx, topk_weights, route_counts, recv_counts, route_slots, arrivals, + recv_x, recv_x_sf, recv_weights, src_ranks, src_tokens, src_topk, + l1_fp8, l1_sf, l2_x, l2_x_sf, barrier, ) if check_capacity: local_max = recv_counts.max() dist.all_reduce(local_max, op=dist.ReduceOp.MAX, group=group) assert local_max.item() <= capacity, f"expert capacity {capacity} is smaller than received routes {local_max.item()}" fused_l2( - l2_x, - l2_fp8, - l2_x_sf, - l2_sf, - recv_counts, - src_ranks, - src_tokens, - src_topk, - combine, - barrier, - out, + l2_x, l2_fp8, l2_x_sf, l2_sf, recv_counts, src_ranks, + src_tokens, src_topk, combine, barrier, out, ) return out @@ -1419,16 +1012,8 @@ def run_pipeline(check_capacity: bool = False): if args.check: expected = torch_reference( - x_fp8_src, - x_sf_src, - topk_idx_src, - topk_weights_src, - l1_fp8_src, - l1_sf_src, - l2_fp8_src, - l2_sf_src, - group, - activation_clamp, + x_fp8_src, x_sf_src, topk_idx_src, topk_weights_src, l1_fp8_src, + l1_sf_src, l2_fp8_src, l2_sf_src, group, activation_clamp, ) diff = calc_diff(actual, expected) assert diff < args.diff_tol, f"rank {local_rank}: diff={diff} exceeds {args.diff_tol}" @@ -1440,13 +1025,7 @@ def run_pipeline(check_capacity: bool = False): for _ in range(args.warmup): run_pipeline() reset_state() - latency = do_bench( - run_pipeline, - warmup=0, - rep=args.rep, - post_fn=reset_state, - group=group, - ) + latency = do_bench(run_pipeline, warmup=0, rep=args.rep, post_fn=reset_state, group=group) if local_rank == 0: print( f"tilescale sm90 fp8 mega moe: model={model_name} family={shape_family} " From add83500c9659cea0005dd6645ede52af9bcf84a Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Thu, 27 Aug 2026 11:39:47 +0800 Subject: [PATCH 21/30] perf(distributed): retune SM90 mega MoE L1 schedule Flash M=8192 on 4x H200 goes from 6152.7us to 4034.5us (-34.4%) and the smoke diff improves an order of magnitude (1.7e-4 -> 1.4e-5). Changes, each measured in isolation: - Flat m-task queue replaces the per-tile expert rescan in both kernels. The old scheduler walked all experts of a wave twice per tile; with experts-per-wave at 64 that dominated everything else (-22%). - `clear_accum=True` lets WGMMA overwrite its accumulator instead of an explicit per-k-step clear of a 64x256 fp32 tile (L2 -8.7%). - Default L1 pipeline stages 5 -> 3. Three ties five at M<=512 and wins 0.6%/2.1% at M=2048/8192, so the extra shared memory is not earning its keep. - FP32 running sum. A BF16 accumulator needs a quarter-rate F2FP per element pair to narrow the WGMMA output every k-step, and costs an order of magnitude of accuracy. - Quantize straight into the shared staging tile, dropping a block_m x block_n/2 fp8 fragment. - A TMA stage now holds num_k_sub contiguous scale-group sub-tiles, so one barrier round-trip covers several of them while WGMMA still consumes them one scale group at a time. The register split is the subtle one. `setmaxnreg` was being dropped outright -- ptxas reported "(C7507) 'setmaxnreg' ignored to maintain minimum register requirements" and SASS contained no USETMAXREG, so the warp-specialised budgets never took effect. The cause is that the dec/inc lived in their own if/else, separate from the specialised code: once that branch rejoins, ptxas cannot prove the deallocating threads never reach the high-pressure path. Nesting dispatch/producer under the dec and the consumer under the inc makes it stick. Budgets are now 64/192, chosen because spilling tracks the frontend budget rather than the math one (40/48/56 give 72/16/0 bytes) and because a split summing to exactly 65536 compiles but deadlocks at run time. Adds --profile-phases plus tuning overrides used to derive the above. Co-Authored-By: Claude Opus 5 (1M context) --- .../mega_moe/example_sm90_fp8_mega_moe.py | 614 +++++++++--------- 1 file changed, 317 insertions(+), 297 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index eb11530027..ee17b7af47 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -76,7 +76,10 @@ def select_manual_warp_configs( routed_tokens = num_tokens * num_topk high_sm = num_sms >= 100 - l1_stages = 5 + # Measured on Flash (4x H200): three stages ties five at M<=512 and wins + # 0.6%/2.1% at M=2048/8192, so the deeper default is not worth its + # shared memory. + l1_stages = 3 l2_stages = 3 generic_experts_per_wave = num_experts_per_rank if num_experts_per_rank <= routed_tokens <= 4 * num_experts_per_rank: @@ -174,6 +177,8 @@ def fused_l1_swiglu_manual_warp_kernel( threads: int = 384, pipeline_stages: int = 5, num_experts_per_wave: int | None = None, + frontend_regs_override: int | None = None, + math_regs_override: int | None = None, ): num_experts_per_rank = num_experts // num_ranks num_experts_per_wave = num_experts_per_wave or num_experts_per_rank @@ -182,10 +187,9 @@ def fused_l1_swiglu_manual_warp_kernel( num_routes = num_tokens * num_topk num_m_blocks = T.ceildiv(capacity, block_m) num_n_blocks = T.ceildiv(l1_n, block_n) - num_expert_waves = num_experts_per_rank // num_experts_per_wave - max_tiles_per_expert_wave = num_experts_per_wave * num_m_blocks * num_n_blocks - max_tile_rounds_per_expert_wave = T.ceildiv(max_tiles_per_expert_wave, num_sms) num_k_blocks = hidden // block_k + num_k_sub = block_k // SCALE_GRANULARITY + assert block_k % SCALE_GRANULARITY == 0 # WG0 is the frontend: warps 0-1 dispatch routes and warps 2-3 issue TMA. # WG1+ are WGMMA consumers; each owns an N fragment of the CTA tile. warp_size = 32 @@ -207,10 +211,15 @@ def fused_l1_swiglu_manual_warp_kernel( num_l1_scale_groups = l1_n // (2 * SCALE_GRANULARITY) tma_block_n = min(block_n, 256) num_tma_n_blocks = block_n // tma_block_n - frontend_registers = 32 if num_math_threads == 512 else 48 - math_registers = 112 if num_math_threads == 512 else 208 + # Budgets must leave the CTA register pool some slack: 128*fe + 256*math + # exactly at 65536 (e.g. 32/240) compiles but deadlocks at run time. + # Spilling tracks the frontend budget, not the math one -- 40/48/56 give + # 72/16/0 bytes of spill -- so keep the frontend at 64 for a spill-free build. + frontend_registers = frontend_regs_override or (32 if num_math_threads == 512 else 64) + math_registers = math_regs_override or (112 if num_math_threads == 512 else 192) dispatch_leader_lane = 0 - route_threads = 256 + route_threads = num_math_threads + assert route_threads % warp_size == 0 @T.prim_func def main( @@ -222,6 +231,8 @@ def main( recv_counts: T.Tensor((num_experts_per_rank,), T.int32), route_slots: T.Tensor((num_tokens, num_topk), T.int32), arrivals: T.Tensor((num_experts_per_rank, T.ceildiv(capacity, block_m)), T.uint32), + m_tasks: T.Tensor((num_experts_per_rank * num_m_blocks,), T.int32), + num_m_tasks: T.Tensor((1,), T.int32), recv_x: T.Tensor((num_experts_per_rank, capacity, hidden), T.float8_e4m3fn), recv_x_sf: T.Tensor((num_experts_per_rank, capacity, num_scale_groups), T.float32), recv_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), @@ -240,18 +251,23 @@ def main( src_rank = T.alloc_local((1,), T.int32) src_rank[0] = T.get_rank() - a_shared = T.alloc_shared((pipeline_stages, block_m, block_k), T.float8_e4m3fn) - b_shared = T.alloc_shared((pipeline_stages, block_n, block_k), T.float8_e4m3fn) + # A TMA stage holds num_k_sub contiguous SCALE_GRANULARITY-deep + # sub-tiles: one barrier round-trip covers all of them, while WGMMA + # still consumes them one scale group at a time. + a_shared = T.alloc_shared((pipeline_stages, num_k_sub, block_m, SCALE_GRANULARITY), T.float8_e4m3fn) + b_shared = T.alloc_shared((pipeline_stages, num_k_sub, block_n, SCALE_GRANULARITY), T.float8_e4m3fn) out_shared = T.alloc_shared((block_m, block_n // 2), T.float8_e4m3fn) stage_barriers = T.alloc_barrier([producer_threads] * pipeline_stages + [num_math_threads] * pipeline_stages) + # Routing runs on the math warpgroups, which hold the large budget. + route_tid = tx - math_begin if bid == 0: - if tx < route_threads: - for reset_expert in T.serial(tx, num_experts, route_threads): + if tx >= math_begin: + for reset_expert in T.serial(route_tid, num_experts, route_threads): route_counts[src_rank[0], reset_expert] = 0 T.sync_threads(7, route_threads) - for assign_route in T.serial(tx, num_routes, route_threads): + for assign_route in T.serial(route_tid, num_routes, route_threads): assign_token = assign_route // num_topk assign_topk = assign_route % num_topk assign_expert = topk_idx[assign_token, assign_topk] @@ -261,22 +277,22 @@ def main( route_slots[assign_token, assign_topk] = -1 T.sync_threads(7, route_threads) - publish_warp = tx // warp_size + publish_warp = route_tid // warp_size for publish_rank in T.serial(publish_warp, num_ranks, route_threads // warp_size): if publish_rank != src_rank[0]: T.put_warp(T.address_of(route_counts[src_rank[0], 0]), T.address_of(route_counts[src_rank[0], 0]), num_experts, dst_pe=publish_rank, unroll_factor=8) T.barrier_blocks(barrier[0]) - if tx < route_threads: - for count_local_expert in T.serial(tx, num_experts_per_rank, route_threads): + if tx >= math_begin: + for count_local_expert in T.serial(route_tid, num_experts_per_rank, route_threads): recv_count = T.alloc_var(T.int32, init=0) recv_expert = src_rank[0] * num_experts_per_rank + count_local_expert for count_rank in T.serial(num_ranks): recv_count += route_counts[count_rank, recv_expert] recv_counts[count_local_expert] = recv_count - for prefix_route in T.serial(tx, num_routes, route_threads): + for prefix_route in T.serial(route_tid, num_routes, route_threads): prefix_token = prefix_route // num_topk prefix_topk = prefix_route % num_topk prefix_expert = topk_idx[prefix_token, prefix_topk] @@ -287,190 +303,151 @@ def main( prefix_slot += route_counts[prefix_rank, prefix_expert] route_slots[prefix_token, prefix_topk] = prefix_slot + T.sync_threads(7, route_threads) + if route_tid == 0: + task_cursor = T.alloc_var(T.int32, init=0) + for task_expert in T.serial(num_experts_per_rank): + task_expert_m_blocks = T.ceildiv(T.min(recv_counts[task_expert], capacity), block_m) + for task_m in T.serial(num_m_blocks): + if task_m < task_expert_m_blocks: + m_tasks[task_cursor] = task_expert * num_m_blocks + task_m + task_cursor += 1 + num_m_tasks[0] = task_cursor + T.sync_grid() + dispatch_warp = tx // warp_size + dispatch_lane = tx % warp_size + # setmaxnreg is only honored when the dec/inc dominates the specialized + # code itself: with a separate if/else, ptxas cannot prove the + # deallocating threads never reach the high-pressure path and drops both. if tx < math_begin: T.dec_max_nreg(frontend_registers) + if tx < dispatch_threads: + # WG0 warps 0-1 publish route metadata and pull remote activation rows. + for metadata_route in T.serial(bid * dispatch_threads + tx, num_routes, num_sms * dispatch_threads): + metadata_token = metadata_route // num_topk + metadata_topk = metadata_route % num_topk + metadata_expert = topk_idx[metadata_token, metadata_topk] + metadata_slot = route_slots[metadata_token, metadata_topk] + if metadata_expert >= 0 and metadata_slot >= 0 and metadata_slot < capacity: + metadata_rank = metadata_expert // num_experts_per_rank + metadata_local_expert = metadata_expert % num_experts_per_rank + T.st(recv_weights[metadata_local_expert, metadata_slot], topk_weights[metadata_token, metadata_topk], dst_pe=metadata_rank) + T.st(src_tokens[metadata_local_expert, metadata_slot], metadata_token, dst_pe=metadata_rank) + T.st(src_topk[metadata_local_expert, metadata_slot], metadata_topk, dst_pe=metadata_rank) + T.st(src_ranks[metadata_local_expert, metadata_slot], src_rank[0], scope="sys", sem="release", dst_pe=metadata_rank) + + for pull_idx in T.serial(bid * dispatch_warps + dispatch_warp, num_m_tasks[0] * block_m, num_sms * dispatch_warps): + pull_m_task = m_tasks[pull_idx // block_m] + pull_expert = pull_m_task // num_m_blocks + pull_slot = (pull_m_task % num_m_blocks) * block_m + pull_idx % block_m + if pull_slot < recv_counts[pull_expert]: + if dispatch_lane == dispatch_leader_lane: + T.wait_ge(src_ranks[pull_expert, pull_slot], 0, scope=T.WaitScope.SYS, semantics=T.WaitSemantics.ACQUIRE) + T.sync_warp() + pull_rank = src_ranks[pull_expert, pull_slot] + pull_token = src_tokens[pull_expert, pull_slot] + T.get_warp(T.address_of(x[pull_token, 0]), T.address_of(recv_x[pull_expert, pull_slot, 0]), hidden, src_pe=pull_rank, unroll_factor=8) + T.get_warp(T.address_of(x_sf[pull_token, 0]), T.address_of(recv_x_sf[pull_expert, pull_slot, 0]), num_scale_groups, src_pe=pull_rank, unroll_factor=8) + T.sync_warp() + if dispatch_lane == dispatch_leader_lane: + T.atom_add(arrivals[pull_expert, pull_slot // block_m], 1, scope="gpu", sem="release") + + if tx >= producer_begin and tx < producer_end: + # WG0 warps 2-3 keep the shared-memory pipeline filled with TMA. + producer_step = T.alloc_var(T.int32, init=0) + for producer_task in T.serial(bid, num_m_tasks[0] * num_n_blocks, num_sms): + producer_n = producer_task % num_n_blocks + producer_m_task = m_tasks[producer_task // num_n_blocks] + producer_m = producer_m_task % num_m_blocks + producer_expert = producer_m_task // num_m_blocks + if producer_n * block_n < l1_n: + producer_arrivals = T.min(block_m, recv_counts[producer_expert] - producer_m * block_m) + if tx == producer_begin: + T.wait_ge(arrivals[producer_expert, producer_m], producer_arrivals, scope=T.WaitScope.GPU, semantics=T.WaitSemantics.ACQUIRE) + T.sync_threads(5, producer_threads) + for producer_k in T.serial(num_k_blocks): + producer_stage = (producer_step + producer_k) % pipeline_stages + producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 + T.mbarrier_wait_parity(stage_barriers[pipeline_stages + producer_stage], producer_phase ^ 1) + for producer_ks in T.unroll(num_k_sub): + producer_sf_k = producer_k * num_k_sub + producer_ks + T.tma_copy( + recv_x[ + producer_expert, + producer_m * block_m : (producer_m + 1) * block_m, + producer_sf_k * SCALE_GRANULARITY : (producer_sf_k + 1) * SCALE_GRANULARITY, + ], + a_shared[producer_stage, producer_ks, :, :], + barrier=stage_barriers[producer_stage], + ) + for producer_n_block in T.serial(num_tma_n_blocks): + T.tma_copy( + l1_weight[ + producer_expert, + producer_n * block_n + + producer_n_block * tma_block_n : producer_n * block_n + + (producer_n_block + 1) * tma_block_n, + producer_sf_k * SCALE_GRANULARITY : (producer_sf_k + 1) * SCALE_GRANULARITY, + ], + b_shared[ + producer_stage, + producer_ks, + producer_n_block + * tma_block_n : (producer_n_block + 1) * tma_block_n, + :, + ], + barrier=stage_barriers[producer_stage], + ) + T.mbarrier_arrive(stage_barriers[producer_stage]) + producer_step += num_k_blocks + else: T.inc_max_nreg(math_registers) - - dispatch_warp = tx // warp_size - dispatch_lane = tx % warp_size - if tx < dispatch_threads: - # WG0 warps 0-1 publish route metadata and pull remote activation rows. - for metadata_route in T.serial(bid * dispatch_threads + tx, num_routes, num_sms * dispatch_threads): - metadata_token = metadata_route // num_topk - metadata_topk = metadata_route % num_topk - metadata_expert = topk_idx[metadata_token, metadata_topk] - metadata_slot = route_slots[metadata_token, metadata_topk] - if metadata_expert >= 0 and metadata_slot >= 0 and metadata_slot < capacity: - metadata_rank = metadata_expert // num_experts_per_rank - metadata_local_expert = metadata_expert % num_experts_per_rank - T.st(recv_weights[metadata_local_expert, metadata_slot], topk_weights[metadata_token, metadata_topk], dst_pe=metadata_rank) - T.st(src_tokens[metadata_local_expert, metadata_slot], metadata_token, dst_pe=metadata_rank) - T.st(src_topk[metadata_local_expert, metadata_slot], metadata_topk, dst_pe=metadata_rank) - T.st(src_ranks[metadata_local_expert, metadata_slot], src_rank[0], scope="sys", sem="release", dst_pe=metadata_rank) - - for pull_idx in T.serial(bid * dispatch_warps + dispatch_warp, num_experts_per_rank * capacity, num_sms * dispatch_warps): - pull_expert = pull_idx // capacity - pull_slot = pull_idx % capacity - if pull_slot < recv_counts[pull_expert]: - if dispatch_lane == dispatch_leader_lane: - T.wait_ge(src_ranks[pull_expert, pull_slot], 0, scope=T.WaitScope.SYS, semantics=T.WaitSemantics.ACQUIRE) - T.sync_warp() - pull_rank = src_ranks[pull_expert, pull_slot] - pull_token = src_tokens[pull_expert, pull_slot] - T.get_warp(T.address_of(x[pull_token, 0]), T.address_of(recv_x[pull_expert, pull_slot, 0]), hidden, src_pe=pull_rank, unroll_factor=8) - T.get_warp(T.address_of(x_sf[pull_token, 0]), T.address_of(recv_x_sf[pull_expert, pull_slot, 0]), num_scale_groups, src_pe=pull_rank, unroll_factor=8) - T.sync_warp() - if dispatch_lane == dispatch_leader_lane: - T.atom_add(arrivals[pull_expert, pull_slot // block_m], 1, scope="gpu", sem="release") - - if tx >= producer_begin and tx < producer_end: - # WG0 warps 2-3 keep the shared-memory pipeline filled with TMA. - producer_step = T.alloc_var(T.int32, init=0) - producer_wave_tile = T.alloc_var(T.int32, init=bid) - for producer_schedule_step in T.serial(num_expert_waves * max_tile_rounds_per_expert_wave): - producer_expert_wave = producer_schedule_step // max_tile_rounds_per_expert_wave - producer_tile_round = producer_schedule_step % max_tile_rounds_per_expert_wave - producer_wave_begin = producer_expert_wave * num_experts_per_wave - producer_wave_num_tiles = T.alloc_var(T.int32, init=0) - for producer_wave_expert_offset in T.serial(num_experts_per_wave): - producer_wave_expert = producer_wave_begin + producer_wave_expert_offset - producer_wave_num_tiles += T.ceildiv(T.min(recv_counts[producer_wave_expert], capacity), block_m) * num_n_blocks - - if producer_wave_tile < producer_wave_num_tiles: - producer_tile_offset = T.alloc_var(T.int32, init=producer_wave_tile) - producer_expert = T.alloc_var(T.int32, init=-1) - producer_m = T.alloc_var(T.int32, init=0) - producer_n = T.alloc_var(T.int32, init=0) - for producer_wave_expert_offset in T.serial(num_experts_per_wave): - producer_wave_expert = producer_wave_begin + producer_wave_expert_offset - producer_expert_m_blocks = T.ceildiv(T.min(recv_counts[producer_wave_expert], capacity), block_m) - producer_expert_tiles = producer_expert_m_blocks * num_n_blocks - if producer_expert < 0: - if producer_tile_offset < producer_expert_tiles: - producer_expert = producer_wave_expert - producer_m = producer_tile_offset // num_n_blocks - producer_n = producer_tile_offset % num_n_blocks - else: - producer_tile_offset -= producer_expert_tiles - - if producer_expert >= 0 and producer_n * block_n < l1_n: - producer_arrivals = T.min(block_m, recv_counts[producer_expert] - producer_m * block_m) - if tx == producer_begin: - T.wait_ge(arrivals[producer_expert, producer_m], producer_arrivals, scope=T.WaitScope.GPU, semantics=T.WaitSemantics.ACQUIRE) - T.sync_threads(5, producer_threads) - for producer_k in T.serial(num_k_blocks): - producer_stage = (producer_step + producer_k) % pipeline_stages - producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 - T.mbarrier_wait_parity(stage_barriers[pipeline_stages + producer_stage], producer_phase ^ 1) - T.tma_copy( - recv_x[ - producer_expert, - producer_m * block_m : (producer_m + 1) * block_m, - producer_k * block_k : (producer_k + 1) * block_k, - ], - a_shared[producer_stage, :, :], - barrier=stage_barriers[producer_stage], - ) - for producer_n_block in T.serial(num_tma_n_blocks): - T.tma_copy( - l1_weight[ - producer_expert, - producer_n * block_n - + producer_n_block * tma_block_n : producer_n * block_n - + (producer_n_block + 1) * tma_block_n, - producer_k * block_k : (producer_k + 1) * block_k, - ], - b_shared[ - producer_stage, - producer_n_block - * tma_block_n : (producer_n_block + 1) * tma_block_n, - :, - ], - barrier=stage_barriers[producer_stage], - ) - T.mbarrier_arrive(stage_barriers[producer_stage]) - producer_step += num_k_blocks - producer_wave_tile += num_sms - - if producer_tile_round == max_tile_rounds_per_expert_wave - 1 and producer_wave_num_tiles > 0: - producer_wave_tile -= producer_wave_num_tiles - - elif tx >= math_begin: # WG1+ consume TMA stages, run WGMMA, and emit quantized L1 rows. partial = T.alloc_fragment((block_m, block_n), T.float32) - accum = T.alloc_fragment((block_m, block_n), T.bfloat16) + # FP32 running sum: a BF16 accumulator would need a quarter-rate + # F2FP per element pair to narrow the WGMMA output every k-step, + # and it costs an order of magnitude of accuracy (1.7e-4 -> 1.4e-5). + accum = T.alloc_fragment((block_m, block_n), T.float32) gate = T.alloc_fragment((block_m, block_n // 2), T.float32) - gate_grouped = T.reshape(gate, (block_m, num_output_scale_groups, SCALE_GRANULARITY)) up = T.alloc_fragment((block_m, block_n // 2), T.float32) + gate_grouped = T.reshape(gate, (block_m, num_output_scale_groups, SCALE_GRANULARITY)) amax = T.alloc_fragment((block_m, num_output_scale_groups), T.float32) scale = T.alloc_fragment((block_m, num_output_scale_groups), T.float32) - quant_fp8 = T.alloc_fragment((block_m, block_n // 2), T.float8_e4m3fn) act_scale = T.alloc_fragment((block_m,), T.float32) weight_scale = T.alloc_local((2 * num_output_scale_groups,), T.float32) consumer_step = T.alloc_var(T.int32, init=0) - consumer_wave_tile = T.alloc_var(T.int32, init=bid) - - for consumer_schedule_step in T.serial(num_expert_waves * max_tile_rounds_per_expert_wave): - consumer_expert_wave = consumer_schedule_step // max_tile_rounds_per_expert_wave - consumer_tile_round = consumer_schedule_step % max_tile_rounds_per_expert_wave - consumer_wave_begin = consumer_expert_wave * num_experts_per_wave - consumer_wave_num_tiles = T.alloc_var(T.int32, init=0) - for consumer_wave_expert_offset in T.serial(num_experts_per_wave): - consumer_wave_expert = consumer_wave_begin + consumer_wave_expert_offset - consumer_wave_num_tiles += T.ceildiv(T.min(recv_counts[consumer_wave_expert], capacity), block_m) * num_n_blocks - - if consumer_wave_tile < consumer_wave_num_tiles: - consumer_tile_offset = T.alloc_var(T.int32, init=consumer_wave_tile) - consumer_expert = T.alloc_var(T.int32, init=-1) - consumer_m = T.alloc_var(T.int32, init=0) - consumer_n = T.alloc_var(T.int32, init=0) - for consumer_wave_expert_offset in T.serial(num_experts_per_wave): - consumer_wave_expert = consumer_wave_begin + consumer_wave_expert_offset - consumer_expert_m_blocks = T.ceildiv(T.min(recv_counts[consumer_wave_expert], capacity), block_m) - consumer_expert_tiles = consumer_expert_m_blocks * num_n_blocks - if consumer_expert < 0: - if consumer_tile_offset < consumer_expert_tiles: - consumer_expert = consumer_wave_expert - consumer_m = consumer_tile_offset // num_n_blocks - consumer_n = consumer_tile_offset % num_n_blocks - else: - consumer_tile_offset -= consumer_expert_tiles - - if consumer_expert >= 0 and consumer_n * block_n < l1_n: + for consumer_task in T.serial(bid, num_m_tasks[0] * num_n_blocks, num_sms): + consumer_n = consumer_task % num_n_blocks + consumer_m_task = m_tasks[consumer_task // num_n_blocks] + consumer_m = consumer_m_task % num_m_blocks + consumer_expert = consumer_m_task // num_m_blocks + if consumer_n * block_n < l1_n: T.clear(partial) T.clear(accum) for consumer_k in T.serial(num_k_blocks): consumer_stage = (consumer_step + consumer_k) % pipeline_stages consumer_phase = ((consumer_step + consumer_k) // pipeline_stages) & 1 T.mbarrier_wait_parity(stage_barriers[consumer_stage], consumer_phase) - T.gemm(a_shared[consumer_stage, :, :], b_shared[consumer_stage, :, :], partial, transpose_B=True) - for i in T.Parallel(block_m): - act_scale[i] = recv_x_sf[ - consumer_expert, - consumer_m * block_m + i, - consumer_k, - ] - for scale_group in T.serial(num_output_scale_groups): - weight_scale[2 * scale_group] = l1_weight_sf[ - consumer_expert, - consumer_n * num_output_scale_groups + scale_group, - consumer_k, - ] - weight_scale[2 * scale_group + 1] = l1_weight_sf[ - consumer_expert, - num_l1_scale_groups - + consumer_n * num_output_scale_groups - + scale_group, - consumer_k, - ] - for i, j in T.Parallel(block_m, block_n): - accum[i, j] = T.cast(partial[i, j], T.bfloat16) * T.cast( - act_scale[i] * weight_scale[2 * (j // (2 * SCALE_GRANULARITY)) + (j % 16) // 8], T.bfloat16 - ) + accum[i, j] - T.clear(partial) + # One TMA stage spans num_k_sub scale groups; WGMMA and promotion + # still run per SCALE_GRANULARITY so the per-128 scales stay exact. + for consumer_ks in T.unroll(num_k_sub): + consumer_sf_k = consumer_k * num_k_sub + consumer_ks + for scale_group in T.serial(num_output_scale_groups): + weight_scale[2 * scale_group] = l1_weight_sf[consumer_expert, consumer_n * num_output_scale_groups + scale_group, consumer_sf_k] + weight_scale[2 * scale_group + 1] = l1_weight_sf[consumer_expert, num_l1_scale_groups + consumer_n * num_output_scale_groups + scale_group, consumer_sf_k] + for i in T.Parallel(block_m): + act_scale[i] = recv_x_sf[consumer_expert, consumer_m * block_m + i, consumer_sf_k] + T.gemm( + a_shared[consumer_stage, consumer_ks, :, :], + b_shared[consumer_stage, consumer_ks, :, :], + partial, transpose_B=True, clear_accum=True) + for i, j in T.Parallel(block_m, block_n): + accum[i, j] = partial[i, j] * ( + act_scale[i] * weight_scale[2 * (j // (2 * SCALE_GRANULARITY)) + (j % 16) // 8] + ) + accum[i, j] T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) consumer_step += num_k_blocks for i, j in T.Parallel(block_m, block_n // 2): @@ -500,8 +477,7 @@ def main( ] = scale[i, scale_group] for i, j in T.Parallel(block_m, block_n // 2): gate[i, j] = T.clamp(gate[i, j] / scale[i, j // SCALE_GRANULARITY], -FP8_MAX, FP8_MAX) - T.copy(gate, quant_fp8) - T.copy(quant_fp8, out_shared) + T.copy(gate, out_shared) T.copy( out_shared, l2_x[ @@ -510,10 +486,6 @@ def main( consumer_n * (block_n // 2), ], ) - consumer_wave_tile += num_sms - - if consumer_tile_round == max_tile_rounds_per_expert_wave - 1 and consumer_wave_num_tiles > 0: - consumer_wave_tile -= consumer_wave_num_tiles return main @@ -531,20 +503,18 @@ def fused_l2_scatter_reduce_manual_warp_kernel( block_m: int = 64, block_n: int = 256, block_k: int = 128, - reduce_block_m: int = 8, - reduce_block_h: int = 128, threads: int = 384, pipeline_stages: int = 3, num_experts_per_wave: int | None = None, + use_put_warp_scatter: bool = False, ): num_experts_per_wave = num_experts_per_wave or num_experts_per_rank assert num_experts_per_rank % num_experts_per_wave == 0 num_m_blocks = T.ceildiv(capacity, block_m) num_n_blocks = T.ceildiv(hidden, block_n) - num_expert_waves = num_experts_per_rank // num_experts_per_wave - max_tiles_per_expert_wave = num_experts_per_wave * num_m_blocks * num_n_blocks - max_tile_rounds_per_expert_wave = T.ceildiv(max_tiles_per_expert_wave, num_sms) num_k_blocks = intermediate_hidden // block_k + reduce_block_m = 8 + reduce_block_h = 128 num_reduce_n_blocks = T.ceildiv(hidden, reduce_block_h) num_reduce_m_blocks = T.ceildiv(num_tokens, reduce_block_m) num_reduce_tiles = num_reduce_n_blocks * num_reduce_m_blocks @@ -576,6 +546,8 @@ def main( a_sf: T.Tensor((num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), T.float32), b_sf: T.Tensor((num_experts_per_rank, hidden // SCALE_GRANULARITY, intermediate_hidden // SCALE_GRANULARITY), T.float32), recv_counts: T.Tensor((num_experts_per_rank,), T.int32), + m_tasks: T.Tensor((num_experts_per_rank * num_m_blocks,), T.int32), + num_m_tasks: T.Tensor((1,), T.int32), src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), src_topk: T.Tensor((num_experts_per_rank, capacity), T.int32), @@ -583,11 +555,13 @@ def main( barrier: T.Tensor((num_ranks,), T.int32), out: T.Tensor((num_tokens, hidden), T.bfloat16), ): + assert capacity > 0 with T.Kernel(num_sms, threads=threads) as bid: tx = T.get_thread_binding() a_shared = T.alloc_shared((pipeline_stages, block_m, block_k), T.float8_e4m3fn) b_shared = T.alloc_shared((pipeline_stages, block_n, block_k), T.float8_e4m3fn) a_sf_shared = T.alloc_shared((pipeline_stages, block_m), T.float32) + scatter_shared = T.alloc_shared((block_m, block_n), T.bfloat16) reduce_shared = T.alloc_shared((reduce_block_m, reduce_block_h), T.bfloat16) stage_barriers = T.alloc_barrier([producer_threads] * pipeline_stages + [num_math_threads] * pipeline_stages) @@ -599,39 +573,16 @@ def main( if tx >= producer_begin and tx < producer_end: # WG0 warps 2-3 keep the L2 TMA stages filled. producer_step = T.alloc_var(T.int32, init=0) - producer_wave_tile = T.alloc_var(T.int32, init=bid) - for producer_schedule_step in T.serial(num_expert_waves * max_tile_rounds_per_expert_wave): - producer_expert_wave = producer_schedule_step // max_tile_rounds_per_expert_wave - producer_tile_round = producer_schedule_step % max_tile_rounds_per_expert_wave - producer_wave_begin = producer_expert_wave * num_experts_per_wave - producer_wave_num_tiles = T.alloc_var(T.int32, init=0) - for producer_wave_expert_offset in T.serial(num_experts_per_wave): - producer_wave_expert = producer_wave_begin + producer_wave_expert_offset - producer_wave_num_tiles += T.ceildiv(T.min(recv_counts[producer_wave_expert], capacity), block_m) * num_n_blocks - - if producer_wave_tile < producer_wave_num_tiles: - producer_tile_offset = T.alloc_var(T.int32, init=producer_wave_tile) - producer_expert = T.alloc_var(T.int32, init=-1) - producer_m = T.alloc_var(T.int32, init=0) - producer_n = T.alloc_var(T.int32, init=0) - for producer_wave_expert_offset in T.serial(num_experts_per_wave): - producer_wave_expert = producer_wave_begin + producer_wave_expert_offset - producer_expert_m_blocks = T.ceildiv(T.min(recv_counts[producer_wave_expert], capacity), block_m) - producer_expert_tiles = producer_expert_m_blocks * num_n_blocks - if producer_expert < 0: - if producer_tile_offset < producer_expert_tiles: - producer_expert = producer_wave_expert - producer_m = producer_tile_offset // num_n_blocks - producer_n = producer_tile_offset % num_n_blocks - else: - producer_tile_offset -= producer_expert_tiles - - if ( - producer_expert >= 0 - and producer_expert < num_experts_per_rank - and producer_n * block_n < hidden - and num_k_blocks * block_k == intermediate_hidden - ): + for producer_task in T.serial(bid, num_m_tasks[0] * num_n_blocks, num_sms): + producer_n = producer_task % num_n_blocks + producer_m_task = m_tasks[producer_task // num_n_blocks] + producer_m = producer_m_task % num_m_blocks + producer_expert = producer_m_task // num_m_blocks + if ( + producer_expert < num_experts_per_rank + and producer_n * block_n < hidden + and num_k_blocks * block_k == intermediate_hidden + ): for producer_k in T.serial(num_k_blocks): producer_stage = (producer_step + producer_k) % pipeline_stages producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 @@ -657,10 +608,6 @@ def main( a_sf_shared[producer_stage, tx - producer_begin] = a_sf[producer_expert, producer_m * block_m + tx - producer_begin, producer_k] T.mbarrier_arrive(stage_barriers[producer_stage]) producer_step += num_k_blocks - producer_wave_tile += num_sms - - if producer_tile_round == max_tile_rounds_per_expert_wave - 1 and producer_wave_num_tiles > 0: - producer_wave_tile -= producer_wave_num_tiles elif tx >= math_begin: # WG1-2 run WGMMA and scatter their BF16 column pairs remotely. @@ -669,50 +616,27 @@ def main( act_scale = T.alloc_fragment((block_m,), T.float32) weight_scale = T.alloc_local((2,), T.float32) consumer_step = T.alloc_var(T.int32, init=0) - consumer_wave_tile = T.alloc_var(T.int32, init=bid) scatter_dst_rank = T.alloc_var(T.int32, init=0) scatter_dst_token = T.alloc_var(T.int32, init=0) scatter_dst_topk = T.alloc_var(T.int32, init=0) - for consumer_schedule_step in T.serial(num_expert_waves * max_tile_rounds_per_expert_wave): - consumer_expert_wave = consumer_schedule_step // max_tile_rounds_per_expert_wave - consumer_tile_round = consumer_schedule_step % max_tile_rounds_per_expert_wave - consumer_wave_begin = consumer_expert_wave * num_experts_per_wave - consumer_wave_num_tiles = T.alloc_var(T.int32, init=0) - for consumer_wave_expert_offset in T.serial(num_experts_per_wave): - consumer_wave_expert = consumer_wave_begin + consumer_wave_expert_offset - consumer_wave_num_tiles += T.ceildiv(T.min(recv_counts[consumer_wave_expert], capacity), block_m) * num_n_blocks - - if consumer_wave_tile < consumer_wave_num_tiles: - consumer_tile_offset = T.alloc_var(T.int32, init=consumer_wave_tile) - consumer_expert = T.alloc_var(T.int32, init=-1) - consumer_m = T.alloc_var(T.int32, init=0) - consumer_n = T.alloc_var(T.int32, init=0) - for consumer_wave_expert_offset in T.serial(num_experts_per_wave): - consumer_wave_expert = consumer_wave_begin + consumer_wave_expert_offset - consumer_expert_m_blocks = T.ceildiv(T.min(recv_counts[consumer_wave_expert], capacity), block_m) - consumer_expert_tiles = consumer_expert_m_blocks * num_n_blocks - if consumer_expert < 0: - if consumer_tile_offset < consumer_expert_tiles: - consumer_expert = consumer_wave_expert - consumer_m = consumer_tile_offset // num_n_blocks - consumer_n = consumer_tile_offset % num_n_blocks - else: - consumer_tile_offset -= consumer_expert_tiles - - if ( - consumer_expert >= 0 - and consumer_expert < num_experts_per_rank - and consumer_n * block_n < hidden - and num_k_blocks * block_k == intermediate_hidden - ): + for consumer_task in T.serial(bid, num_m_tasks[0] * num_n_blocks, num_sms): + consumer_n = consumer_task % num_n_blocks + consumer_m_task = m_tasks[consumer_task // num_n_blocks] + consumer_m = consumer_m_task % num_m_blocks + consumer_expert = consumer_m_task // num_m_blocks + if ( + consumer_expert < num_experts_per_rank + and consumer_n * block_n < hidden + and num_k_blocks * block_k == intermediate_hidden + ): T.clear(partial) T.clear(accum) for consumer_k in T.serial(num_k_blocks): consumer_stage = (consumer_step + consumer_k) % pipeline_stages consumer_phase = ((consumer_step + consumer_k) // pipeline_stages) & 1 T.mbarrier_wait_parity(stage_barriers[consumer_stage], consumer_phase) - T.gemm(a_shared[consumer_stage, :, :], b_shared[consumer_stage, :, :], partial, transpose_B=True) + T.gemm(a_shared[consumer_stage, :, :], b_shared[consumer_stage, :, :], partial, transpose_B=True, clear_accum=True) for i in T.Parallel(block_m): act_scale[i] = a_sf_shared[consumer_stage, i] weight_scale[0] = b_sf[consumer_expert, consumer_n * 2, consumer_k] @@ -721,40 +645,68 @@ def main( accum[i, j] = T.cast(partial[i, j], T.bfloat16) * T.cast( act_scale[i] * weight_scale[j // 128], T.bfloat16 ) + accum[i, j] - T.clear(partial) T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) consumer_step += num_k_blocks - # Map the SM90 M64 WGMMA fragment owners directly to - # packed BF16 remote stores, without a shared-memory epilogue. - scatter_math_thread = tx - math_begin - scatter_wg = scatter_math_thread // warpgroup_size - scatter_warp_in_wg = (scatter_math_thread % warpgroup_size) // warp_size - scatter_lane = scatter_math_thread % warp_size - for scatter_row_half in T.serial(2): - row = scatter_warp_in_wg * 16 + scatter_row_half * 8 + scatter_lane // 4 - pool_row = consumer_m * block_m + row - if pool_row < recv_counts[consumer_expert]: - scatter_dst_rank = src_ranks[consumer_expert, pool_row] - scatter_dst_token = src_tokens[consumer_expert, pool_row] - scatter_dst_topk = src_topk[consumer_expert, pool_row] - if ( - scatter_dst_rank >= 0 - and scatter_dst_rank < num_ranks - and scatter_dst_token >= 0 - and scatter_dst_token < num_tokens - and scatter_dst_topk >= 0 - and scatter_dst_topk < num_topk - ): - for scatter_col_chunk in T.serial(block_n // math_warpgroups // 8): - scatter_col = scatter_wg * (block_n // math_warpgroups) + scatter_col_chunk * 8 + (scatter_lane % 4) * 2 - scatter_value_lo = T.alloc_var(T.uint16, init=T.reinterpret(accum[row, scatter_col], T.uint16)) - scatter_value_hi = T.alloc_var(T.uint16, init=T.reinterpret(accum[row, scatter_col + 1], T.uint16)) - scatter_value = T.alloc_var(T.uint32, init=T.cast(scatter_value_lo, T.uint32) | (T.cast(scatter_value_hi, T.uint32) << 16)) - T.st(combine[scatter_dst_token, scatter_dst_topk, consumer_n * block_n + scatter_col], scatter_value, dst_pe=scatter_dst_rank) - consumer_wave_tile += num_sms - - if consumer_tile_round == max_tile_rounds_per_expert_wave - 1 and consumer_wave_num_tiles > 0: - consumer_wave_tile -= consumer_wave_num_tiles + if use_put_warp_scatter: + # Stage the complete tile once, then let each math warp + # scatter eight rows with aligned 16-byte remote stores. + T.copy(accum, scatter_shared) + T.sync_threads(4, num_math_threads) + scatter_warp = (tx - math_begin) // warp_size + for row_in_warp in T.serial(block_m // (num_math_threads // warp_size)): + row = scatter_warp * (block_m // (num_math_threads // warp_size)) + row_in_warp + pool_row = consumer_m * block_m + row + if pool_row < recv_counts[consumer_expert]: + if tx % warp_size == 0: + scatter_dst_rank = src_ranks[consumer_expert, pool_row] + scatter_dst_token = src_tokens[consumer_expert, pool_row] + scatter_dst_topk = src_topk[consumer_expert, pool_row] + dst_rank = T.shfl_sync(scatter_dst_rank, 0) + dst_token = T.shfl_sync(scatter_dst_token, 0) + dst_topk = T.shfl_sync(scatter_dst_topk, 0) + if ( + dst_rank >= 0 + and dst_rank < num_ranks + and dst_token >= 0 + and dst_token < num_tokens + and dst_topk >= 0 + and dst_topk < num_topk + ): + T.put_warp( + T.address_of(scatter_shared[row, 0]), + T.address_of(combine[dst_token, dst_topk, consumer_n * block_n]), + block_n, + dst_pe=dst_rank, + unroll_factor=1, + ) + T.sync_threads(4, num_math_threads) + else: + # Direct scatter maps fragment owners to packed BF16 stores. + scatter_math_thread = tx - math_begin + scatter_wg = scatter_math_thread // warpgroup_size + scatter_warp_in_wg = (scatter_math_thread % warpgroup_size) // warp_size + scatter_lane = scatter_math_thread % warp_size + for scatter_row_half in T.serial(2): + row = scatter_warp_in_wg * 16 + scatter_row_half * 8 + scatter_lane // 4 + pool_row = consumer_m * block_m + row + if pool_row < recv_counts[consumer_expert]: + scatter_dst_rank = src_ranks[consumer_expert, pool_row] + scatter_dst_token = src_tokens[consumer_expert, pool_row] + scatter_dst_topk = src_topk[consumer_expert, pool_row] + if ( + scatter_dst_rank >= 0 + and scatter_dst_rank < num_ranks + and scatter_dst_token >= 0 + and scatter_dst_token < num_tokens + and scatter_dst_topk >= 0 + and scatter_dst_topk < num_topk + ): + for scatter_col_chunk in T.serial(block_n // math_warpgroups // 8): + scatter_col = scatter_wg * (block_n // math_warpgroups) + scatter_col_chunk * 8 + (scatter_lane % 4) * 2 + scatter_value_lo = T.alloc_var(T.uint16, init=T.reinterpret(accum[row, scatter_col], T.uint16)) + scatter_value_hi = T.alloc_var(T.uint16, init=T.reinterpret(accum[row, scatter_col + 1], T.uint16)) + scatter_value = T.alloc_var(T.uint32, init=T.cast(scatter_value_lo, T.uint32) | (T.cast(scatter_value_hi, T.uint32) << 16)) + T.st(combine[scatter_dst_token, scatter_dst_topk, consumer_n * block_n + scatter_col], scatter_value, dst_pe=scatter_dst_rank) T.fence_sys() T.sync_grid() @@ -899,7 +851,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): rank, num_ranks, group = init_dist(local_rank, num_local_ranks) assert rank == local_rank and num_ranks == num_local_ranks - num_sms = torch.cuda.get_device_properties(local_rank).multi_processor_count + num_sms = args.num_sms or torch.cuda.get_device_properties(local_rank).multi_processor_count allocator = get_allocator( size=_allocator_size_bytes(num_tokens, hidden, intermediate_hidden, num_experts_per_rank, num_topk, capacity), device=f"cuda:{local_rank}", @@ -911,18 +863,27 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): ) shape_family, l1_config, l2_config = select_manual_warp_configs(hidden, intermediate_hidden, num_tokens, num_topk, num_experts_per_rank, num_sms) + if args.l1_block_k is not None: + l1_config["block_k"] = args.l1_block_k + if args.l1_stages is not None: + l1_config["pipeline_stages"] = args.l1_stages for phase, config in (("l1", l1_config), ("l2", l2_config)): requested = getattr(args, f"{phase}_experts_per_wave", None) if requested is not None: assert requested > 0 and num_experts_per_rank % requested == 0 config["num_experts_per_wave"] = requested + l2_scatter = getattr(args, "l2_scatter", "auto") + use_put_warp_scatter = l2_scatter == "warp" or (l2_scatter == "auto" and num_tokens >= 256) kernel_specs = [ fused_l1_swiglu_manual_warp_kernel( num_tokens, hidden, 2 * intermediate_hidden, num_experts, num_topk, num_ranks, capacity, num_sms, - activation_clamp=activation_clamp, **l1_config, + activation_clamp=activation_clamp, + frontend_regs_override=args.l1_frontend_regs, + math_regs_override=args.l1_math_regs, **l1_config, ), fused_l2_scatter_reduce_manual_warp_kernel( num_tokens, hidden, intermediate_hidden, num_experts_per_rank, num_topk, num_ranks, capacity, num_sms, **l2_config, + use_put_warp_scatter=use_put_warp_scatter, ), ] kernels = [tilelang.compile(spec, compile_once=True, compile_group=group) for spec in kernel_specs] @@ -971,6 +932,8 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): src_topk = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) route_slots = allocator_tensor((num_tokens, num_topk), torch.int32, allocator=allocator) arrivals = allocator_tensor((num_experts_per_rank, (capacity + 63) // 64), torch.uint32, allocator=allocator) + m_tasks = allocator_tensor((num_experts_per_rank * ((capacity + 63) // 64),), torch.int32, allocator=allocator) + num_m_tasks = allocator_tensor((1,), torch.int32, allocator=allocator) l2_x = allocator_tensor((num_experts_per_rank, capacity, intermediate_hidden), torch.float8_e4m3fn, allocator=allocator) l2_x_sf = allocator_tensor((num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), torch.float32, allocator=allocator) combine = allocator_tensor((num_tokens, num_topk, hidden), torch.bfloat16, allocator=allocator) @@ -992,7 +955,7 @@ def reset_state(): def run_pipeline(check_capacity: bool = False): fused_l1( x, x_sf, topk_idx, topk_weights, route_counts, recv_counts, route_slots, arrivals, - recv_x, recv_x_sf, recv_weights, src_ranks, src_tokens, src_topk, + m_tasks, num_m_tasks, recv_x, recv_x_sf, recv_weights, src_ranks, src_tokens, src_topk, l1_fp8, l1_sf, l2_x, l2_x_sf, barrier, ) if check_capacity: @@ -1000,7 +963,7 @@ def run_pipeline(check_capacity: bool = False): dist.all_reduce(local_max, op=dist.ReduceOp.MAX, group=group) assert local_max.item() <= capacity, f"expert capacity {capacity} is smaller than received routes {local_max.item()}" fused_l2( - l2_x, l2_fp8, l2_x_sf, l2_sf, recv_counts, src_ranks, + l2_x, l2_fp8, l2_x_sf, l2_sf, recv_counts, m_tasks, num_m_tasks, src_ranks, src_tokens, src_topk, combine, barrier, out, ) return out @@ -1032,9 +995,59 @@ def run_pipeline(check_capacity: bool = False): f"M={num_tokens} H={hidden} IH={intermediate_hidden} E={num_experts} " f"topk={num_topk} capacity={capacity} " f"epw={l1_config['num_experts_per_wave']}/{l2_config['num_experts_per_wave']} " + f"l2_scatter={'warp' if use_put_warp_scatter else 'direct'} " f"latency={latency * 1000:.1f} us" ) + if args.profile_phases > 0: + reset_state() + for _ in range(args.warmup): + run_pipeline() + reset_state() + + samples = [] + for _ in range(args.profile_phases): + dist.barrier(group=group) + events = [torch.cuda.Event(enable_timing=True) for _ in range(3)] + events[0].record() + fused_l1( + x, x_sf, topk_idx, topk_weights, route_counts, recv_counts, route_slots, arrivals, + m_tasks, num_m_tasks, recv_x, recv_x_sf, recv_weights, src_ranks, src_tokens, src_topk, + l1_fp8, l1_sf, l2_x, l2_x_sf, barrier, + ) + events[1].record() + fused_l2( + l2_x, l2_fp8, l2_x_sf, l2_sf, recv_counts, m_tasks, num_m_tasks, src_ranks, + src_tokens, src_topk, combine, barrier, out, + ) + events[2].record() + events[2].synchronize() + local = torch.tensor( + [events[0].elapsed_time(events[1]), events[1].elapsed_time(events[2])], + dtype=torch.float32, + device="cuda", + ) + gathered = [torch.empty_like(local) for _ in range(num_ranks)] + dist.all_gather(gathered, local, group=group) + if local_rank == 0: + samples.append(torch.stack(gathered).cpu()) + reset_state() + + if local_rank == 0: + stacked = torch.stack(samples) + max_rank_median = stacked.max(dim=1).values.median(dim=0).values * 1000 + rank_medians = stacked.median(dim=0).values * 1000 + print( + f"phase profile: samples={args.profile_phases} max-rank median " + f"l1={max_rank_median[0]:.1f} us l2={max_rank_median[1]:.1f} us " + f"total={max_rank_median[0] + max_rank_median[1]:.1f} us" + ) + for phase_idx, phase_name in enumerate(("l1", "l2")): + rank_values = ", ".join( + f"r{r}={rank_medians[r, phase_idx]:.1f}" for r in range(num_ranks) + ) + print(f"phase profile {phase_name} rank medians (us): {rank_values}") + allocator.close() dist.destroy_process_group() @@ -1051,11 +1064,18 @@ def run_pipeline(check_capacity: bool = False): parser.add_argument("--capacity", type=int, default=None) parser.add_argument("--l1-experts-per-wave", type=int, default=None) parser.add_argument("--l2-experts-per-wave", type=int, default=None) + parser.add_argument("--l2-scatter", choices=("auto", "direct", "warp"), default="auto") parser.add_argument("--activation-clamp", type=float, default=10.0) parser.add_argument("--seed", type=int, default=0) parser.add_argument("--diff-tol", type=float, default=0.01) parser.add_argument("--warmup", type=int, default=1) parser.add_argument("--rep", type=int, default=1) + parser.add_argument("--profile-phases", type=int, default=0) + parser.add_argument("--num-sms", type=int, default=None) + parser.add_argument("--l1-block-k", type=int, default=None) + parser.add_argument("--l1-frontend-regs", type=int, default=None) + parser.add_argument("--l1-math-regs", type=int, default=None) + parser.add_argument("--l1-stages", type=int, default=None) parser.add_argument("--check", action="store_true") parser.add_argument("--print-source", action="store_true") args = parser.parse_args() From 68088edcc9653ddaf1ab9c8ca0d5c09b85b84444 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Thu, 27 Aug 2026 11:48:01 +0800 Subject: [PATCH 22/30] perf(distributed): stage mega MoE activation scales in shared memory ncu points at two LDG instructions carrying 96% of the L1 kernel's excessive sectors. They are the per-token activation scale loads: with `recv_x_sf` laid out (expert, token, group), reading a column walks a 128-byte stride, and the thread mapping has four lanes share an address so each warpgroup re-reads the same 64 floats -- 8x redundant across the CTA. The producer now reads that column once into shared and the math warpgroups read shared instead. Placement matters more than the staging itself: issuing the strided load at the top of the k-step, before the TMA copies, lets its latency hide behind TMA issue. Doing the same load just before `mbarrier_arrive` instead measured 1% *slower* than not staging at all, because it serialises the load into the producer's critical path. Flash M=8192 on 4x H200: 4034.5us -> 3984.6us. Four samples, no overlap between the two configurations. Also tried and rejected: transposing `recv_x_sf` to scale-group major so the consumer read is contiguous and TMA-eligible (TMA needs the innermost dim to be a multiple of 16 bytes, which a single float is not). That is what PR383 does -- its SFA descriptor is MN-major and the scales arrive by TMA on the same stage barrier as A/B. Here it cost 11.7% on L1: the dispatch gathers by token while the GEMM consumes by k, so a transposed pool turns one vectorised `get_warp` into a per-lane remote scalar read plus a strided local write, and that loss exceeds the read-side gain. Co-Authored-By: Claude Opus 5 (1M context) --- .../mega_moe/example_sm90_fp8_mega_moe.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index ee17b7af47..d63f198da0 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -257,6 +257,12 @@ def main( a_shared = T.alloc_shared((pipeline_stages, num_k_sub, block_m, SCALE_GRANULARITY), T.float8_e4m3fn) b_shared = T.alloc_shared((pipeline_stages, num_k_sub, block_n, SCALE_GRANULARITY), T.float8_e4m3fn) out_shared = T.alloc_shared((block_m, block_n // 2), T.float8_e4m3fn) + # Stage the per-token activation scales in shared memory: read from + # global in the math warpgroups, they cost an 8x-redundant, + # 128-byte-strided LDG that ncu flags as the dominant uncoalesced + # access. The producer issues its load before the TMAs so the + # latency hides behind TMA issue rather than delaying the arrive. + act_sf_shared = T.alloc_shared((pipeline_stages, block_m), T.float32) stage_barriers = T.alloc_barrier([producer_threads] * pipeline_stages + [num_math_threads] * pipeline_stages) # Routing runs on the math warpgroups, which hold the large budget. @@ -371,6 +377,9 @@ def main( producer_stage = (producer_step + producer_k) % pipeline_stages producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 T.mbarrier_wait_parity(stage_barriers[pipeline_stages + producer_stage], producer_phase ^ 1) + producer_sf = T.alloc_local((1,), T.float32) + producer_sf[0] = recv_x_sf[ + producer_expert, producer_m * block_m + tx - producer_begin, producer_k] for producer_ks in T.unroll(num_k_sub): producer_sf_k = producer_k * num_k_sub + producer_ks T.tma_copy( @@ -400,6 +409,7 @@ def main( ], barrier=stage_barriers[producer_stage], ) + act_sf_shared[producer_stage, tx - producer_begin] = producer_sf[0] T.mbarrier_arrive(stage_barriers[producer_stage]) producer_step += num_k_blocks @@ -439,7 +449,7 @@ def main( weight_scale[2 * scale_group] = l1_weight_sf[consumer_expert, consumer_n * num_output_scale_groups + scale_group, consumer_sf_k] weight_scale[2 * scale_group + 1] = l1_weight_sf[consumer_expert, num_l1_scale_groups + consumer_n * num_output_scale_groups + scale_group, consumer_sf_k] for i in T.Parallel(block_m): - act_scale[i] = recv_x_sf[consumer_expert, consumer_m * block_m + i, consumer_sf_k] + act_scale[i] = act_sf_shared[consumer_stage, i] T.gemm( a_shared[consumer_stage, consumer_ks, :, :], b_shared[consumer_stage, consumer_ks, :, :], From 4068861448423d4a2fdab1904e0201bc08c91759 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Thu, 27 Aug 2026 12:00:18 +0800 Subject: [PATCH 23/30] perf(distributed): make the mega MoE scale pool scale-group major `recv_x_sf` is now (expert, group, capacity), so the producer reads a whole block_m column of activation scales contiguously instead of walking a 128-byte stride. An earlier attempt at this layout cost 11.7% on L1, but it changed two things at once: it also replaced the vectorised `get_warp` pull with a per-lane remote scalar load. Splitting them shows the transpose itself is cheap and the pull primitive was the whole regression. The pull now keeps `get_warp` -- reading the remote row contiguously into a small shared staging tile -- and only the local write scatters, one scale group per lane. Flash M=8192 on 4x H200: 3978.9us -> 3930.2us. This is the layout half of what PR383 gets from its MN-major SFA descriptor. The other half, feeding the scales in by TMA on the stage barrier, needs the innermost dimension to be a multiple of 16 bytes; a block_m column of floats now qualifies, so it is worth revisiting. Co-Authored-By: Claude Opus 5 (1M context) --- .../mega_moe/example_sm90_fp8_mega_moe.py | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index d63f198da0..868a675527 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -234,7 +234,7 @@ def main( m_tasks: T.Tensor((num_experts_per_rank * num_m_blocks,), T.int32), num_m_tasks: T.Tensor((1,), T.int32), recv_x: T.Tensor((num_experts_per_rank, capacity, hidden), T.float8_e4m3fn), - recv_x_sf: T.Tensor((num_experts_per_rank, capacity, num_scale_groups), T.float32), + recv_x_sf: T.Tensor((num_experts_per_rank, num_scale_groups, capacity), T.float32), recv_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), @@ -263,6 +263,7 @@ def main( # access. The producer issues its load before the TMAs so the # latency hides behind TMA issue rather than delaying the arrive. act_sf_shared = T.alloc_shared((pipeline_stages, block_m), T.float32) + pull_sf_shared = T.alloc_shared((dispatch_warps, num_scale_groups), T.float32) stage_barriers = T.alloc_barrier([producer_threads] * pipeline_stages + [num_math_threads] * pipeline_stages) # Routing runs on the math warpgroups, which hold the large budget. @@ -355,7 +356,12 @@ def main( pull_rank = src_ranks[pull_expert, pull_slot] pull_token = src_tokens[pull_expert, pull_slot] T.get_warp(T.address_of(x[pull_token, 0]), T.address_of(recv_x[pull_expert, pull_slot, 0]), hidden, src_pe=pull_rank, unroll_factor=8) - T.get_warp(T.address_of(x_sf[pull_token, 0]), T.address_of(recv_x_sf[pull_expert, pull_slot, 0]), num_scale_groups, src_pe=pull_rank, unroll_factor=8) + # Keep the remote read contiguous, then scatter into the scale-group + # major pool so the producer can read a block_m column contiguously. + T.get_warp(T.address_of(x_sf[pull_token, 0]), T.address_of(pull_sf_shared[dispatch_warp, 0]), num_scale_groups, src_pe=pull_rank, unroll_factor=8) + T.sync_warp() + for pull_sf_group in T.serial(dispatch_lane, num_scale_groups, warp_size): + recv_x_sf[pull_expert, pull_sf_group, pull_slot] = pull_sf_shared[dispatch_warp, pull_sf_group] T.sync_warp() if dispatch_lane == dispatch_leader_lane: T.atom_add(arrivals[pull_expert, pull_slot // block_m], 1, scope="gpu", sem="release") @@ -379,7 +385,7 @@ def main( T.mbarrier_wait_parity(stage_barriers[pipeline_stages + producer_stage], producer_phase ^ 1) producer_sf = T.alloc_local((1,), T.float32) producer_sf[0] = recv_x_sf[ - producer_expert, producer_m * block_m + tx - producer_begin, producer_k] + producer_expert, producer_k, producer_m * block_m + tx - producer_begin] for producer_ks in T.unroll(num_k_sub): producer_sf_k = producer_k * num_k_sub + producer_ks T.tma_copy( @@ -935,7 +941,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): route_counts = allocator_tensor((num_ranks, num_experts), torch.int32, allocator=allocator) recv_counts = allocator_tensor((num_experts_per_rank,), torch.int32, allocator=allocator) recv_x = allocator_tensor((num_experts_per_rank, capacity, hidden), torch.float8_e4m3fn, allocator=allocator) - recv_x_sf = allocator_tensor((num_experts_per_rank, capacity, hidden // SCALE_GRANULARITY), torch.float32, allocator=allocator) + recv_x_sf = allocator_tensor((num_experts_per_rank, hidden // SCALE_GRANULARITY, capacity), torch.float32, allocator=allocator) recv_weights = allocator_tensor((num_experts_per_rank, capacity), torch.float32, allocator=allocator) src_ranks = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) src_tokens = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) From c8e2f787d465dbe3874a3f4771eec12b90589761 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Thu, 27 Aug 2026 14:02:22 +0800 Subject: [PATCH 24/30] perf(distributed): apply the L1 scale treatment to L2 `l2_x_sf` is now scale-group major, matching `recv_x_sf`, and the L2 producer issues its scale read at the top of the k step instead of immediately before `mbarrier_arrive` -- the same placement that was worth 2.2% on L1. L2 goes from 1769.6us to 1750.7us (two samples each, both lower). The win is smaller than L1's because L2 already staged its scales in shared memory and never had the 8x redundancy; only the contiguous read and the earlier issue are new here. Flash M=8192 end to end: 3927.9 -> 3925.5us. Also tried and rejected: dropping the tail `sync_threads` after the put_warp scatter, on the theory that the next iteration's fragment-to-shared copy already carries the loop-carried hazard barrier. The two configurations' samples overlap completely (1738.5/1757.5 vs 1737.7/1748.7). An earlier note measured 0.94% for this, but that was on the wave scheduler, where fewer, larger tiles made the per-tile barrier a bigger share; the flat task queue has already absorbed it. Co-Authored-By: Claude Opus 5 (1M context) --- .../mega_moe/example_sm90_fp8_mega_moe.py | 14 +- ...example_sm90_fp8_mega_moe_single_kernel.py | 1327 +++++++++++++++++ 2 files changed, 1336 insertions(+), 5 deletions(-) create mode 100644 examples/distributed/mega_moe/example_sm90_fp8_mega_moe_single_kernel.py diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 868a675527..6484b7eef4 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -242,7 +242,7 @@ def main( l1_weight: T.Tensor((num_experts_per_rank, l1_n, hidden), T.float8_e4m3fn), l1_weight_sf: T.Tensor((num_experts_per_rank, l1_n // SCALE_GRANULARITY, hidden // SCALE_GRANULARITY), T.float32), l2_x: T.Tensor((num_experts_per_rank, capacity, l1_n // 2), T.float8_e4m3fn), - l2_x_sf: T.Tensor((num_experts_per_rank, capacity, l1_n // (2 * SCALE_GRANULARITY)), T.float32), + l2_x_sf: T.Tensor((num_experts_per_rank, l1_n // (2 * SCALE_GRANULARITY), capacity), T.float32), barrier: T.Tensor((num_ranks,), T.int32), ): T.annotate_pass_configs({tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) @@ -488,8 +488,8 @@ def main( scale[i, scale_group] = T.max(amax[i, scale_group], 1e-4) / FP8_MAX l2_x_sf[ consumer_expert, - consumer_m * block_m + i, consumer_n * num_output_scale_groups + scale_group, + consumer_m * block_m + i, ] = scale[i, scale_group] for i, j in T.Parallel(block_m, block_n // 2): gate[i, j] = T.clamp(gate[i, j] / scale[i, j // SCALE_GRANULARITY], -FP8_MAX, FP8_MAX) @@ -559,7 +559,7 @@ def fused_l2_scatter_reduce_manual_warp_kernel( def main( a: T.Tensor((num_experts_per_rank, capacity, intermediate_hidden), T.float8_e4m3fn), b: T.Tensor((num_experts_per_rank, hidden, intermediate_hidden), T.float8_e4m3fn), - a_sf: T.Tensor((num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), T.float32), + a_sf: T.Tensor((num_experts_per_rank, intermediate_hidden // SCALE_GRANULARITY, capacity), T.float32), b_sf: T.Tensor((num_experts_per_rank, hidden // SCALE_GRANULARITY, intermediate_hidden // SCALE_GRANULARITY), T.float32), recv_counts: T.Tensor((num_experts_per_rank,), T.int32), m_tasks: T.Tensor((num_experts_per_rank * num_m_blocks,), T.int32), @@ -589,6 +589,7 @@ def main( if tx >= producer_begin and tx < producer_end: # WG0 warps 2-3 keep the L2 TMA stages filled. producer_step = T.alloc_var(T.int32, init=0) + producer_sf = T.alloc_local((1,), T.float32) for producer_task in T.serial(bid, num_m_tasks[0] * num_n_blocks, num_sms): producer_n = producer_task % num_n_blocks producer_m_task = m_tasks[producer_task // num_n_blocks] @@ -603,6 +604,9 @@ def main( producer_stage = (producer_step + producer_k) % pipeline_stages producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 T.mbarrier_wait_parity(stage_barriers[pipeline_stages + producer_stage], producer_phase ^ 1) + # Issue the scale read before the TMAs so its latency hides behind TMA + # issue rather than delaying the arrive (measured both ways on L1). + producer_sf[0] = a_sf[producer_expert, producer_k, producer_m * block_m + tx - producer_begin] T.tma_copy( a[ producer_expert, @@ -621,7 +625,7 @@ def main( b_shared[producer_stage, :, :], barrier=stage_barriers[producer_stage], ) - a_sf_shared[producer_stage, tx - producer_begin] = a_sf[producer_expert, producer_m * block_m + tx - producer_begin, producer_k] + a_sf_shared[producer_stage, tx - producer_begin] = producer_sf[0] T.mbarrier_arrive(stage_barriers[producer_stage]) producer_step += num_k_blocks @@ -951,7 +955,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): m_tasks = allocator_tensor((num_experts_per_rank * ((capacity + 63) // 64),), torch.int32, allocator=allocator) num_m_tasks = allocator_tensor((1,), torch.int32, allocator=allocator) l2_x = allocator_tensor((num_experts_per_rank, capacity, intermediate_hidden), torch.float8_e4m3fn, allocator=allocator) - l2_x_sf = allocator_tensor((num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), torch.float32, allocator=allocator) + l2_x_sf = allocator_tensor((num_experts_per_rank, intermediate_hidden // SCALE_GRANULARITY, capacity), torch.float32, allocator=allocator) combine = allocator_tensor((num_tokens, num_topk, hidden), torch.bfloat16, allocator=allocator) out = allocator_tensor((num_tokens, hidden), torch.bfloat16, allocator=allocator) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe_single_kernel.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe_single_kernel.py new file mode 100644 index 0000000000..c46096c6fb --- /dev/null +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe_single_kernel.py @@ -0,0 +1,1327 @@ +"""Experimental single-kernel SM90 FP8 Mega MoE using TileScale. + +Dedicated dispatch, GEMM, and combine warps form a cross-wave pipeline. L1 +epilogues publish per-(expert, M-block) readiness for L2, while rank-wave +completion flags let combine consume wave N as the GEMM warps enter wave N+1. +The established two-kernel example remains the stable baseline. +""" + +from __future__ import annotations + +import argparse +import os +from typing import Tuple + +import torch +import torch.distributed as dist +import torch.multiprocessing + +import tilelang +import tilelang.language as T +from tilelang.distributed.allocator import get_allocator +from tilelang.distributed.bench import do_bench +from tilelang.distributed.host import init_dist + +import example_sm90_fp8_mega_moe as two_kernel + + +os.environ.setdefault("NCCL_DEBUG", "ERROR") + +MODEL_CONFIGS = two_kernel.MODEL_CONFIGS +FP8_MAX = two_kernel.FP8_MAX +SCALE_GRANULARITY = two_kernel.SCALE_GRANULARITY + + +def select_single_kernel_config( + hidden: int, + intermediate_hidden: int, + num_tokens: int, + num_topk: int, + num_experts_per_rank: int, + num_sms: int, +) -> Tuple[str, dict[str, int]]: + """Collapse the validated L1/L2 schedules into one compatible schedule.""" + family, l1, l2 = two_kernel.select_manual_warp_configs( + hidden, + intermediate_hidden, + num_tokens, + num_topk, + num_experts_per_rank, + num_sms, + ) + for key in ("block_m", "block_n", "block_k", "threads"): + assert l1[key] == l2[key] + preferred_wave_size = 16 if family == "compact" else 24 + preferred_wave_size = min(preferred_wave_size, num_experts_per_rank) + wave_size = next( + size + for size in range(preferred_wave_size, num_experts_per_rank + 1) + if num_experts_per_rank % size == 0 + ) + return family, { + "block_m": l1["block_m"], + "block_n": l1["block_n"], + "block_k": l1["block_k"], + "threads": 512, + # The fused allocation includes both epilogues. Four stages keeps its + # shared-memory footprint below the SM90 per-CTA limit. + "pipeline_stages": min(max(l1["pipeline_stages"], l2["pipeline_stages"]), 4), + "num_experts_per_wave": wave_size, + } + + +def fused_single_kernel( + num_tokens: int, + hidden: int, + intermediate_hidden: int, + num_experts: int, + num_topk: int, + num_ranks: int, + capacity: int, + num_sms: int, + activation_clamp: float = 10.0, + block_m: int = 64, + block_n: int = 256, + block_k: int = 128, + threads: int = 512, + pipeline_stages: int = 3, + num_experts_per_wave: int | None = None, +): + num_experts_per_rank = num_experts // num_ranks + num_experts_per_wave = num_experts_per_wave or num_experts_per_rank + assert num_experts_per_rank % num_experts_per_wave == 0 + assert block_m == 64 and block_n == 256 and block_k == 128 + assert threads == 512 and pipeline_stages <= 4 + assert hidden % block_n == 0 and intermediate_hidden % block_k == 0 + + l1_n = 2 * intermediate_hidden + num_scale_groups = hidden // SCALE_GRANULARITY + num_routes = num_tokens * num_topk + num_m_blocks = T.ceildiv(capacity, block_m) + l1_num_n_blocks = l1_n // block_n + l2_num_n_blocks = hidden // block_n + l1_num_k_blocks = hidden // block_k + l2_num_k_blocks = intermediate_hidden // block_k + num_expert_waves = num_experts_per_rank // num_experts_per_wave + # Task cursors flatten (expert, M block, N block) so L1 and L2 can advance + # independently once their block-level readiness checks are enabled. + l1_total_tasks = num_experts_per_rank * num_m_blocks * l1_num_n_blocks + l2_total_tasks = num_experts_per_rank * num_m_blocks * l2_num_n_blocks + l1_total_rounds = T.ceildiv(l1_total_tasks, num_sms) + l2_total_rounds = T.ceildiv(l2_total_tasks, num_sms) + + num_output_scale_groups = block_n // (2 * SCALE_GRANULARITY) + num_l1_scale_groups = l1_n // (2 * SCALE_GRANULARITY) + combine_block_n = 128 + num_combine_n_blocks = hidden // combine_block_n + # Eight chunks amortize each wave wait while keeping enough tasks to fill + # the four combine warps on all SMs. + combine_n_blocks_per_task = 8 + num_combine_groups = T.ceildiv( + num_combine_n_blocks, combine_n_blocks_per_task + ) + combine_values_per_lane = combine_block_n // 32 + + warp_size = 32 + dispatch_threads = 64 + producer_begin = dispatch_threads + producer_threads = 64 + producer_end = producer_begin + producer_threads + math_begin = producer_end + num_math_threads = 256 + math_end = math_begin + num_math_threads + # Keep the validated producer/math warp IDs unchanged and append combine. + combine_begin = math_end + combine_threads = 128 + combine_end = combine_begin + combine_threads + num_combine_warps = combine_threads // warp_size + math_warps = num_math_threads // warp_size + rows_per_math_warp = block_m // math_warps + route_threads = 256 + + @T.prim_func + def main( + x: T.Tensor((num_tokens, hidden), T.float8_e4m3fn), + x_sf: T.Tensor((num_tokens, num_scale_groups), T.float32), + topk_idx: T.Tensor((num_tokens, num_topk), T.int32), + topk_weights: T.Tensor((num_tokens, num_topk), T.float32), + route_counts: T.Tensor((num_ranks, num_experts), T.int32), + recv_counts: T.Tensor((num_experts_per_rank,), T.int32), + route_slots: T.Tensor((num_tokens, num_topk), T.int32), + dispatch_arrivals: T.Tensor( + (num_experts_per_rank, T.ceildiv(capacity, block_m)), T.uint32 + ), + l2_arrivals: T.Tensor( + (num_experts_per_rank, T.ceildiv(capacity, block_m)), T.uint32 + ), + l2_task_ready: T.Tensor( + (num_experts_per_rank, T.ceildiv(capacity, block_m), l2_num_n_blocks), T.uint32 + ), + recv_x: T.Tensor( + (num_experts_per_rank, capacity, hidden), T.float8_e4m3fn + ), + recv_x_sf: T.Tensor( + (num_experts_per_rank, capacity, num_scale_groups), T.float32 + ), + recv_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), + src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), + src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), + src_topk: T.Tensor((num_experts_per_rank, capacity), T.int32), + l1_weight: T.Tensor( + (num_experts_per_rank, 2 * intermediate_hidden, hidden), + T.float8_e4m3fn, + ), + l1_weight_sf: T.Tensor( + ( + num_experts_per_rank, + 2 * intermediate_hidden // SCALE_GRANULARITY, + hidden // SCALE_GRANULARITY, + ), + T.float32, + ), + l2_weight: T.Tensor( + (num_experts_per_rank, hidden, intermediate_hidden), + T.float8_e4m3fn, + ), + l2_weight_sf: T.Tensor( + ( + num_experts_per_rank, + hidden // SCALE_GRANULARITY, + intermediate_hidden // SCALE_GRANULARITY, + ), + T.float32, + ), + l2_x: T.Tensor( + (num_experts_per_rank, capacity, intermediate_hidden), + T.float8_e4m3fn, + ), + l2_x_sf: T.Tensor( + ( + num_experts_per_rank, + capacity, + intermediate_hidden // SCALE_GRANULARITY, + ), + T.float32, + ), + combine: T.Tensor((num_tokens, num_topk, hidden), T.bfloat16), + barrier: T.Tensor((num_ranks,), T.int32), + out: T.Tensor((num_tokens, hidden), T.bfloat16), + ): + T.annotate_pass_configs({tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) + with T.Kernel(num_sms, threads=threads) as bid: + tx = T.get_thread_binding() + src_rank = T.alloc_local((1,), T.int32) + src_rank[0] = T.get_rank() + + a_shared = T.alloc_shared( + (pipeline_stages, block_m, block_k), T.float8_e4m3fn + ) + b_shared = T.alloc_shared( + (pipeline_stages, block_n, block_k), T.float8_e4m3fn + ) + l2_a_sf_shared = T.alloc_shared( + (pipeline_stages, block_m), T.float32 + ) + l1_out_shared = T.alloc_shared( + (block_m, block_n // 2), T.float8_e4m3fn + ) + l2_out_shared = T.alloc_shared((block_m, block_n), T.bfloat16) + stage_barriers = T.alloc_barrier( + [producer_threads] * pipeline_stages + + [num_math_threads] * pipeline_stages + ) + + if bid == 0: + if tx < route_threads: + for reset_wave in T.serial(T.ceildiv(num_experts, route_threads)): + reset_expert = tx + reset_wave * route_threads + if reset_expert < num_experts: + route_counts[src_rank[0], reset_expert] = 0 + T.sync_threads(7, route_threads) + + for assign_wave in T.serial(T.ceildiv(num_routes, route_threads)): + assign_route = tx + assign_wave * route_threads + if assign_route < num_routes: + assign_token = assign_route // num_topk + assign_topk = assign_route % num_topk + assign_expert = topk_idx[assign_token, assign_topk] + if assign_expert >= 0 and assign_expert < num_experts: + route_slots[assign_token, assign_topk] = T.atomic_add( + route_counts[src_rank[0], assign_expert], + 1, + memory_order="relaxed", + return_prev=True, + ) + else: + route_slots[assign_token, assign_topk] = -1 + T.sync_threads(7, route_threads) + + for publish_wave in T.serial( + T.ceildiv(num_experts * num_ranks, route_threads) + ): + publish_idx = tx + publish_wave * route_threads + if publish_idx < num_experts * num_ranks: + publish_rank = publish_idx // num_experts + publish_expert = publish_idx % num_experts + if publish_rank != src_rank[0]: + T.st( + route_counts[src_rank[0], publish_expert], + route_counts[src_rank[0], publish_expert], + dst_pe=publish_rank, + ) + + T.barrier_blocks(barrier[0]) + + if tx < route_threads: + for count_wave in T.serial( + T.ceildiv(num_experts_per_rank, route_threads) + ): + local_expert = tx + count_wave * route_threads + if local_expert < num_experts_per_rank: + recv_count = T.alloc_var(T.int32, init=0) + recv_expert = ( + src_rank[0] * num_experts_per_rank + local_expert + ) + for count_rank in T.serial(num_ranks): + recv_count += route_counts[count_rank, recv_expert] + recv_counts[local_expert] = recv_count + + for prefix_wave in T.serial(T.ceildiv(num_routes, route_threads)): + prefix_route = tx + prefix_wave * route_threads + if prefix_route < num_routes: + prefix_token = prefix_route // num_topk + prefix_topk = prefix_route % num_topk + prefix_expert = topk_idx[prefix_token, prefix_topk] + prefix_slot = T.alloc_var( + T.int32, + init=route_slots[prefix_token, prefix_topk], + ) + if ( + prefix_expert >= 0 + and prefix_expert < num_experts + and prefix_slot >= 0 + ): + for prefix_rank in T.serial(num_ranks): + if prefix_rank < src_rank[0]: + prefix_slot += route_counts[ + prefix_rank, prefix_expert + ] + route_slots[prefix_token, prefix_topk] = prefix_slot + + T.sync_grid() + + if tx < math_begin or tx >= math_end: + T.dec_max_nreg(48) + else: + T.inc_max_nreg(208) + + dispatch_warp = tx // warp_size + dispatch_lane = tx % warp_size + if tx < dispatch_threads: + for metadata_wave in T.serial( + T.ceildiv(num_routes, num_sms * dispatch_threads) + ): + metadata_route = ( + bid * dispatch_threads + + tx + + metadata_wave * num_sms * dispatch_threads + ) + if metadata_route < num_routes: + metadata_token = metadata_route // num_topk + metadata_topk = metadata_route % num_topk + metadata_expert = topk_idx[metadata_token, metadata_topk] + metadata_slot = route_slots[metadata_token, metadata_topk] + if ( + metadata_expert >= 0 + and metadata_slot >= 0 + and metadata_slot < capacity + ): + metadata_rank = ( + metadata_expert // num_experts_per_rank + ) + metadata_local_expert = ( + metadata_expert % num_experts_per_rank + ) + T.st( + recv_weights[ + metadata_local_expert, metadata_slot + ], + topk_weights[metadata_token, metadata_topk], + dst_pe=metadata_rank, + ) + T.st( + src_tokens[metadata_local_expert, metadata_slot], + metadata_token, + dst_pe=metadata_rank, + ) + T.st( + src_topk[metadata_local_expert, metadata_slot], + metadata_topk, + dst_pe=metadata_rank, + ) + T.st( + src_ranks[metadata_local_expert, metadata_slot], + src_rank[0], + scope="sys", + sem="release", + dst_pe=metadata_rank, + ) + + for pull_wave in T.serial( + T.ceildiv( + num_experts_per_rank * capacity, + num_sms * (dispatch_threads // warp_size), + ) + ): + pull_idx = ( + bid * (dispatch_threads // warp_size) + + dispatch_warp + + pull_wave + * num_sms + * (dispatch_threads // warp_size) + ) + pull_expert = pull_idx // capacity + pull_slot = pull_idx % capacity + if ( + pull_expert < num_experts_per_rank + and pull_slot < recv_counts[pull_expert] + ): + if dispatch_lane == 0: + T.wait_ge( + src_ranks[pull_expert, pull_slot], + 0, + scope=T.WaitScope.SYS, + semantics=T.WaitSemantics.ACQUIRE, + ) + T.sync_warp() + pull_rank = src_ranks[pull_expert, pull_slot] + pull_token = src_tokens[pull_expert, pull_slot] + T.get_warp( + T.address_of(x[pull_token, 0]), + T.address_of(recv_x[pull_expert, pull_slot, 0]), + hidden, + src_pe=pull_rank, + unroll_factor=8, + ) + T.get_warp( + T.address_of(x_sf[pull_token, 0]), + T.address_of(recv_x_sf[pull_expert, pull_slot, 0]), + num_scale_groups, + src_pe=pull_rank, + unroll_factor=8, + ) + T.sync_warp() + if dispatch_lane == 0: + T.atom_add( + dispatch_arrivals[ + pull_expert, pull_slot // block_m + ], + 1, + scope="gpu", + sem="release", + ) + + if tx >= combine_begin and tx < combine_end: + combine_warp = (tx - combine_begin) // warp_size + combine_lane = (tx - combine_begin) % warp_size + combine_accum = T.alloc_local( + ( + combine_n_blocks_per_task, + combine_values_per_lane, + ), + T.float32, + ) + num_combine_tasks = num_tokens * num_combine_groups + for combine_round in T.serial( + T.ceildiv( + num_combine_tasks, + num_sms * num_combine_warps, + ) + ): + combine_task = ( + bid * num_combine_warps + + combine_warp + + combine_round * num_sms * num_combine_warps + ) + if combine_task < num_combine_tasks: + combine_token = combine_task // num_combine_groups + combine_group = combine_task % num_combine_groups + for group_block in T.unroll( + combine_n_blocks_per_task + ): + for value_idx in T.unroll( + combine_values_per_lane + ): + combine_accum[group_block, value_idx] = 0.0 + + for expert_wave in T.serial(num_expert_waves): + wave_begin = expert_wave * num_experts_per_wave + wave_end = wave_begin + num_experts_per_wave + for topk_slot in T.serial(num_topk): + route_expert = topk_idx[ + combine_token, topk_slot + ] + route_slot = route_slots[ + combine_token, topk_slot + ] + route_local_expert = ( + route_expert % num_experts_per_rank + ) + route_rank = ( + route_expert // num_experts_per_rank + ) + if ( + route_expert >= 0 + and route_slot >= 0 + and route_slot < capacity + and route_local_expert >= wave_begin + and route_local_expert < wave_end + ): + for group_block in T.unroll( + combine_n_blocks_per_task + ): + combine_n_block = ( + combine_group + * combine_n_blocks_per_task + + group_block + ) + if combine_n_block < num_combine_n_blocks: + if combine_lane == 0: + T.wait_ge( + l2_task_ready[ + route_local_expert, + route_slot // block_m, + combine_n_block + // (block_n // combine_block_n), + ], + 1, + scope=T.WaitScope.SYS, + semantics=T.WaitSemantics.ACQUIRE, + ) + T.sync_warp() + for value_idx in T.unroll( + combine_values_per_lane + ): + combine_col = ( + combine_n_block + * combine_block_n + + value_idx * warp_size + + combine_lane + ) + combine_accum[ + group_block, value_idx + ] += T.cast( + combine[ + combine_token, + topk_slot, + combine_col, + ], + T.float32, + ) + + for group_block in T.unroll( + combine_n_blocks_per_task + ): + combine_n_block = ( + combine_group * combine_n_blocks_per_task + + group_block + ) + if combine_n_block < num_combine_n_blocks: + for value_idx in T.unroll( + combine_values_per_lane + ): + combine_col = ( + combine_n_block * combine_block_n + + value_idx * warp_size + + combine_lane + ) + out[combine_token, combine_col] = T.cast( + combine_accum[ + group_block, value_idx + ], + T.bfloat16, + ) + + if tx >= producer_begin and tx < producer_end: + producer_step = T.alloc_var(T.int32, init=0) + l1_task = T.alloc_var(T.int32, init=bid) + l2_task = T.alloc_var(T.int32, init=bid) + + # SM90 keeps one producer warpgroup and advances each phase's + # flattened task cursor independently. The L2 cursor waits on + # the per-expert/M-block L1 readiness counter below. + for l1_round in T.serial(l1_total_rounds): + if l1_task < l1_total_tasks: + l1_task_offset = T.alloc_var(T.int32, init=l1_task) + l1_expert = l1_task_offset // (num_m_blocks * l1_num_n_blocks) + l1_task_offset -= l1_expert * num_m_blocks * l1_num_n_blocks + l1_m_block = l1_task_offset // l1_num_n_blocks + l1_n_block = l1_task_offset % l1_num_n_blocks + l1_valid_m_blocks = T.ceildiv( + T.min(recv_counts[l1_expert], capacity), block_m + ) + if l1_m_block < l1_valid_m_blocks: + valid_rows = T.min( + block_m, + recv_counts[l1_expert] - l1_m_block * block_m, + ) + if tx == producer_begin: + T.wait_ge( + dispatch_arrivals[l1_expert, l1_m_block], + valid_rows, + scope=T.WaitScope.GPU, + semantics=T.WaitSemantics.ACQUIRE, + ) + T.sync_threads(5, producer_threads) + for k_block in T.serial(l1_num_k_blocks): + stage = (producer_step + k_block) % pipeline_stages + phase = ((producer_step + k_block) // pipeline_stages) & 1 + T.mbarrier_wait_parity( + stage_barriers[pipeline_stages + stage], phase ^ 1 + ) + T.tma_copy( + recv_x[ + l1_expert, + l1_m_block * block_m : (l1_m_block + 1) * block_m, + k_block * block_k : (k_block + 1) * block_k, + ], + a_shared[stage, :, :], + barrier=stage_barriers[stage], + ) + T.tma_copy( + l1_weight[ + l1_expert, + l1_n_block * block_n : (l1_n_block + 1) * block_n, + k_block * block_k : (k_block + 1) * block_k, + ], + b_shared[stage, :, :], + barrier=stage_barriers[stage], + ) + T.mbarrier_arrive(stage_barriers[stage]) + producer_step += l1_num_k_blocks + l1_task += num_sms + + for l2_round in T.serial(l2_total_rounds): + if l2_task < l2_total_tasks: + l2_task_offset = T.alloc_var(T.int32, init=l2_task) + l2_expert = l2_task_offset // (num_m_blocks * l2_num_n_blocks) + l2_task_offset -= l2_expert * num_m_blocks * l2_num_n_blocks + l2_m_block = l2_task_offset // l2_num_n_blocks + l2_n_block = l2_task_offset % l2_num_n_blocks + l2_valid_m_blocks = T.ceildiv( + T.min(recv_counts[l2_expert], capacity), block_m + ) + if l2_m_block < l2_valid_m_blocks: + if tx == producer_begin: + T.wait_ge( + l2_arrivals[l2_expert, l2_m_block], + l1_num_n_blocks, + scope=T.WaitScope.GPU, + semantics=T.WaitSemantics.ACQUIRE, + ) + T.sync_threads(6, producer_threads) + for k_block in T.serial(l2_num_k_blocks): + stage = (producer_step + k_block) % pipeline_stages + phase = ((producer_step + k_block) // pipeline_stages) & 1 + T.mbarrier_wait_parity( + stage_barriers[pipeline_stages + stage], phase ^ 1 + ) + T.tma_copy( + l2_x[ + l2_expert, + l2_m_block * block_m : (l2_m_block + 1) * block_m, + k_block * block_k : (k_block + 1) * block_k, + ], + a_shared[stage, :, :], + barrier=stage_barriers[stage], + ) + T.tma_copy( + l2_weight[ + l2_expert, + l2_n_block * block_n : (l2_n_block + 1) * block_n, + k_block * block_k : (k_block + 1) * block_k, + ], + b_shared[stage, :, :], + barrier=stage_barriers[stage], + ) + l2_a_sf_shared[stage, tx - producer_begin] = l2_x_sf[ + l2_expert, + l2_m_block * block_m + tx - producer_begin, + k_block, + ] + T.mbarrier_arrive(stage_barriers[stage]) + producer_step += l2_num_k_blocks + l2_task += num_sms + + elif tx >= math_begin and tx < math_end: + partial = T.alloc_fragment((block_m, block_n), T.float32) + accum = T.alloc_fragment((block_m, block_n), T.bfloat16) + gate = T.alloc_fragment((block_m, block_n // 2), T.float32) + gate_grouped = T.reshape( + gate, + ( + block_m, + num_output_scale_groups, + SCALE_GRANULARITY, + ), + ) + up = T.alloc_fragment((block_m, block_n // 2), T.float32) + amax = T.alloc_fragment( + (block_m, num_output_scale_groups), T.float32 + ) + scale = T.alloc_fragment( + (block_m, num_output_scale_groups), T.float32 + ) + quant_fp8 = T.alloc_fragment( + (block_m, block_n // 2), T.float8_e4m3fn + ) + act_scale = T.alloc_fragment((block_m,), T.float32) + weight_scale = T.alloc_local( + (2 * num_output_scale_groups,), T.float32 + ) + consumer_step = T.alloc_var(T.int32, init=0) + l1_task = T.alloc_var(T.int32, init=bid) + l2_task = T.alloc_var(T.int32, init=bid) + scatter_dst_rank = T.alloc_var(T.int32, init=0) + scatter_dst_token = T.alloc_var(T.int32, init=0) + scatter_dst_topk = T.alloc_var(T.int32, init=0) + + # Keep one surrounding scope while task-level completion is + # refined to block-level flags. + for expert_wave in T.serial(1): + + for l1_round in T.serial(l1_total_rounds): + if l1_task < l1_total_tasks: + tile_offset = T.alloc_var(T.int32, init=l1_task) + expert = tile_offset // (num_m_blocks * l1_num_n_blocks) + tile_offset -= expert * num_m_blocks * l1_num_n_blocks + m_block = tile_offset // l1_num_n_blocks + n_block = tile_offset % l1_num_n_blocks + l1_valid_m_blocks = T.ceildiv( + T.min(recv_counts[expert], capacity), block_m + ) + if m_block < l1_valid_m_blocks and ( + expert >= 0 + and l2_num_k_blocks * block_k + == intermediate_hidden + ): + T.clear(partial) + T.clear(accum) + for k_block in T.serial(l1_num_k_blocks): + stage = (consumer_step + k_block) % pipeline_stages + phase = ( + (consumer_step + k_block) + // pipeline_stages + ) & 1 + T.mbarrier_wait_parity( + stage_barriers[stage], phase + ) + T.gemm( + a_shared[stage, :, :], + b_shared[stage, :, :], + partial, + transpose_B=True, + ) + for i in T.Parallel(block_m): + act_scale[i] = recv_x_sf[ + expert, + m_block * block_m + i, + k_block, + ] + for scale_group in T.serial( + num_output_scale_groups + ): + weight_scale[2 * scale_group] = ( + l1_weight_sf[ + expert, + n_block + * num_output_scale_groups + + scale_group, + k_block, + ] + ) + weight_scale[2 * scale_group + 1] = ( + l1_weight_sf[ + expert, + num_l1_scale_groups + + n_block + * num_output_scale_groups + + scale_group, + k_block, + ] + ) + for i, j in T.Parallel(block_m, block_n): + accum[i, j] = ( + T.cast(partial[i, j], T.bfloat16) + * T.cast( + act_scale[i] + * weight_scale[ + 2 + * ( + j + // ( + 2 + * SCALE_GRANULARITY + ) + ) + + (j % 16) // 8 + ], + T.bfloat16, + ) + + accum[i, j] + ) + T.clear(partial) + T.mbarrier_arrive( + stage_barriers[ + pipeline_stages + stage + ] + ) + consumer_step += l1_num_k_blocks + + for i, j in T.Parallel( + block_m, block_n // 2 + ): + gate[i, j] = accum[ + i, (j // 8) * 16 + j % 8 + ] + for i, j in T.Parallel( + block_m, block_n // 2 + ): + up[i, j] = accum[ + i, (j // 8) * 16 + j % 8 + 8 + ] + for i, j in T.Parallel( + block_m, block_n // 2 + ): + clamped_gate = T.min( + gate[i, j], activation_clamp + ) + gate[i, j] = ( + clamped_gate + * T.sigmoid(clamped_gate) + * T.max( + T.min(up[i, j], activation_clamp), + -activation_clamp, + ) + * recv_weights[ + expert, m_block * block_m + i + ] + ) + T.reduce_absmax(gate_grouped, amax, dim=2) + for i, scale_group in T.Parallel( + block_m, num_output_scale_groups + ): + scale[i, scale_group] = ( + T.max(amax[i, scale_group], 1e-4) + / FP8_MAX + ) + l2_x_sf[ + expert, + m_block * block_m + i, + n_block + * num_output_scale_groups + + scale_group, + ] = scale[i, scale_group] + for i, j in T.Parallel( + block_m, block_n // 2 + ): + gate[i, j] = T.clamp( + gate[i, j] + / scale[ + i, j // SCALE_GRANULARITY + ], + -FP8_MAX, + FP8_MAX, + ) + T.copy(gate, quant_fp8) + T.copy(quant_fp8, l1_out_shared) + T.copy( + l1_out_shared, + l2_x[ + expert, + m_block + * block_m : (m_block + 1) + * block_m, + n_block + * (block_n // 2) : (n_block + 1) + * (block_n // 2), + ], + ) + if tx == math_begin: + T.atom_add( + l2_arrivals[expert, m_block], + 1, + scope="gpu", + sem="release", + ) + l1_task += num_sms + + for l2_round in T.serial(l2_total_rounds): + if l2_task < l2_total_tasks: + tile_offset = T.alloc_var(T.int32, init=l2_task) + expert = tile_offset // (num_m_blocks * l2_num_n_blocks) + tile_offset -= expert * num_m_blocks * l2_num_n_blocks + m_block = tile_offset // l2_num_n_blocks + n_block = tile_offset % l2_num_n_blocks + l2_valid_m_blocks = T.ceildiv( + T.min(recv_counts[expert], capacity), block_m + ) + if m_block < l2_valid_m_blocks and expert >= 0: + T.clear(partial) + T.clear(accum) + for k_block in T.serial(l2_num_k_blocks): + stage = (consumer_step + k_block) % pipeline_stages + phase = ( + (consumer_step + k_block) + // pipeline_stages + ) & 1 + T.mbarrier_wait_parity( + stage_barriers[stage], phase + ) + T.gemm( + a_shared[stage, :, :], + b_shared[stage, :, :], + partial, + transpose_B=True, + ) + for i in T.Parallel(block_m): + act_scale[i] = l2_a_sf_shared[ + stage, i + ] + weight_scale[0] = l2_weight_sf[ + expert, n_block * 2, k_block + ] + weight_scale[1] = l2_weight_sf[ + expert, n_block * 2 + 1, k_block + ] + for i, j in T.Parallel(block_m, block_n): + accum[i, j] = ( + T.cast(partial[i, j], T.bfloat16) + * T.cast( + act_scale[i] + * weight_scale[j // 128], + T.bfloat16, + ) + + accum[i, j] + ) + T.clear(partial) + T.mbarrier_arrive( + stage_barriers[ + pipeline_stages + stage + ] + ) + consumer_step += l2_num_k_blocks + T.copy(accum, l2_out_shared) + + scatter_warp = ( + tx - math_begin + ) // warp_size + for row_in_warp in T.serial(rows_per_math_warp): + row = ( + scatter_warp * rows_per_math_warp + + row_in_warp + ) + pool_row = m_block * block_m + row + if pool_row < recv_counts[expert]: + if tx % warp_size == 0: + scatter_dst_rank = src_ranks[ + expert, pool_row + ] + scatter_dst_token = src_tokens[ + expert, pool_row + ] + scatter_dst_topk = src_topk[ + expert, pool_row + ] + dst_rank = T.shfl_sync( + scatter_dst_rank, 0 + ) + dst_token = T.shfl_sync( + scatter_dst_token, 0 + ) + dst_topk = T.shfl_sync( + scatter_dst_topk, 0 + ) + if ( + dst_rank >= 0 + and dst_rank < num_ranks + and dst_token >= 0 + and dst_token < num_tokens + and dst_topk >= 0 + and dst_topk < num_topk + ): + T.put_warp( + T.address_of( + l2_out_shared[row, 0] + ), + T.address_of( + combine[ + dst_token, + dst_topk, + n_block * block_n, + ] + ), + block_n, + dst_pe=dst_rank, + unroll_factor=1, + ) + T.fence_sys() + T.sync_threads(4, num_math_threads) + if tx == math_begin: + for ready_rank in T.serial(num_ranks): + T.st( + l2_task_ready[expert, m_block, n_block], + 1, + scope="sys", + sem="release", + dst_pe=ready_rank, + ) + T.sync_threads(4, num_math_threads) + l2_task += num_sms + + + T.fence_sys() + T.sync_grid() + + return main + + +def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): + model_name, model = two_kernel.resolve_model_config(args) + hidden = model["hidden"] + intermediate_hidden = model["intermediate_hidden"] + num_experts = model["num_experts"] + num_topk = model["num_topk"] + num_tokens = args.num_tokens + activation_clamp = args.activation_clamp + + assert num_tokens > 0 + assert hidden >= 512 and hidden % 256 == 0 + assert intermediate_hidden > 0 and intermediate_hidden % 128 == 0 + assert num_experts > 0 and num_experts % num_local_ranks == 0 + assert 0 < num_topk <= min(32, num_experts) + num_experts_per_rank = num_experts // num_local_ranks + average_recv = ( + num_tokens * num_local_ranks * num_topk + num_experts - 1 + ) // num_experts + capacity = ( + args.capacity + if args.capacity is not None + else (max(average_recv * 2, 64) + 63) // 64 * 64 + ) + assert capacity >= 64 and capacity % 64 == 0 + + rank, num_ranks, group = init_dist(local_rank, num_local_ranks) + assert rank == local_rank and num_ranks == num_local_ranks + num_sms = torch.cuda.get_device_properties(local_rank).multi_processor_count + allocator = get_allocator( + size=two_kernel._allocator_size_bytes( + num_tokens, + hidden, + intermediate_hidden, + num_experts_per_rank, + num_topk, + capacity, + ), + device=f"cuda:{local_rank}", + is_distributed=True, + local_rank=local_rank, + num_local_ranks=num_local_ranks, + group=group, + use_vmm=True, + ) + + shape_family, config = select_single_kernel_config( + hidden, + intermediate_hidden, + num_tokens, + num_topk, + num_experts_per_rank, + num_sms, + ) + if args.experts_per_wave is not None: + assert args.experts_per_wave > 0 + assert num_experts_per_rank % args.experts_per_wave == 0 + config["num_experts_per_wave"] = args.experts_per_wave + if args.pipeline_stages is not None: + assert 2 <= args.pipeline_stages <= 4 + config["pipeline_stages"] = args.pipeline_stages + num_expert_waves = ( + num_experts_per_rank // config["num_experts_per_wave"] + ) + + spec = fused_single_kernel( + num_tokens, + hidden, + intermediate_hidden, + num_experts, + num_topk, + num_ranks, + capacity, + num_sms, + activation_clamp=activation_clamp, + **config, + ) + kernel = tilelang.compile(spec, compile_once=True, compile_group=group) + kernel.initialize(allocator=allocator) + if local_rank == 0 and args.print_source: + print(kernel.get_kernel_source()) + + torch.manual_seed(args.seed + local_rank) + x_bf16 = torch.randn( + (num_tokens, hidden), dtype=torch.bfloat16, device="cuda" + ) + x_fp8_src, x_sf_src = two_kernel.per_token_cast_to_fp8(x_bf16) + scores = torch.randn( + (num_tokens, num_experts), dtype=torch.float32, device="cuda" + ) + topk_weights_src, topk_idx_src = torch.topk( + scores, num_topk, dim=-1, sorted=False + ) + topk_idx_src = topk_idx_src.to(torch.int32) + + l1_bf16 = ( + torch.randn( + (num_experts_per_rank, 2 * intermediate_hidden, hidden), + dtype=torch.bfloat16, + device="cuda", + ) + * 0.05 + ) + l2_bf16 = ( + torch.randn( + (num_experts_per_rank, hidden, intermediate_hidden), + dtype=torch.bfloat16, + device="cuda", + ) + * 0.05 + ) + l1_fp8_src, l1_sf_src = two_kernel.block_cast_to_fp8(l1_bf16) + l2_fp8_src, l2_sf_src = two_kernel.block_cast_to_fp8(l2_bf16) + del scores, l1_bf16, l2_bf16 + + tensor = two_kernel.allocator_tensor + barrier = tensor((num_ranks,), torch.int32, allocator=allocator) + x = tensor(x_fp8_src.shape, x_fp8_src.dtype, allocator=allocator).copy_( + x_fp8_src + ) + x_sf = tensor(x_sf_src.shape, x_sf_src.dtype, allocator=allocator).copy_( + x_sf_src + ) + topk_idx = tensor( + topk_idx_src.shape, topk_idx_src.dtype, allocator=allocator + ).copy_(topk_idx_src) + topk_weights = tensor( + topk_weights_src.shape, topk_weights_src.dtype, allocator=allocator + ).copy_(topk_weights_src) + l1_fp8_kernel = two_kernel.interleave_gate_up_weights(l1_fp8_src) + l1_fp8 = tensor( + l1_fp8_kernel.shape, l1_fp8_kernel.dtype, allocator=allocator + ).copy_(l1_fp8_kernel) + del l1_fp8_kernel + l1_sf = tensor(l1_sf_src.shape, l1_sf_src.dtype, allocator=allocator).copy_( + l1_sf_src + ) + l2_fp8 = tensor( + l2_fp8_src.shape, l2_fp8_src.dtype, allocator=allocator + ).copy_(l2_fp8_src) + l2_sf = tensor(l2_sf_src.shape, l2_sf_src.dtype, allocator=allocator).copy_( + l2_sf_src + ) + + route_counts = tensor( + (num_ranks, num_experts), torch.int32, allocator=allocator + ) + recv_counts = tensor( + (num_experts_per_rank,), torch.int32, allocator=allocator + ) + recv_x = tensor( + (num_experts_per_rank, capacity, hidden), + torch.float8_e4m3fn, + allocator=allocator, + ) + recv_x_sf = tensor( + (num_experts_per_rank, capacity, hidden // SCALE_GRANULARITY), + torch.float32, + allocator=allocator, + ) + recv_weights = tensor( + (num_experts_per_rank, capacity), torch.float32, allocator=allocator + ) + src_ranks = tensor( + (num_experts_per_rank, capacity), torch.int32, allocator=allocator + ) + src_tokens = tensor( + (num_experts_per_rank, capacity), torch.int32, allocator=allocator + ) + src_topk = tensor( + (num_experts_per_rank, capacity), torch.int32, allocator=allocator + ) + route_slots = tensor( + (num_tokens, num_topk), torch.int32, allocator=allocator + ) + dispatch_arrivals = tensor( + (num_experts_per_rank, capacity // 64), + torch.uint32, + allocator=allocator, + ) + l2_arrivals = tensor( + (num_experts_per_rank, capacity // 64), + torch.uint32, + allocator=allocator, + ) + l2_task_ready = tensor( + ( + num_experts_per_rank, + capacity // config["block_m"], + hidden // config["block_n"], + ), + torch.uint32, + allocator=allocator, + ) + l2_x = tensor( + (num_experts_per_rank, capacity, intermediate_hidden), + torch.float8_e4m3fn, + allocator=allocator, + ) + l2_x_sf = tensor( + ( + num_experts_per_rank, + capacity, + intermediate_hidden // SCALE_GRANULARITY, + ), + torch.float32, + allocator=allocator, + ) + combine = tensor( + (num_tokens, num_topk, hidden), torch.bfloat16, allocator=allocator + ) + out = tensor((num_tokens, hidden), torch.bfloat16, allocator=allocator) + + def reset_state(): + route_counts.zero_() + barrier.zero_() + recv_counts.zero_() + dispatch_arrivals.zero_() + l2_arrivals.zero_() + l2_task_ready.zero_() + recv_x.zero_() + recv_x_sf.zero_() + recv_weights.zero_() + src_ranks.fill_(-1) + combine.zero_() + torch.cuda.synchronize() + dist.barrier(group=group) + + def run_pipeline(check_capacity: bool = False): + kernel( + x, + x_sf, + topk_idx, + topk_weights, + route_counts, + recv_counts, + route_slots, + dispatch_arrivals, + l2_arrivals, + l2_task_ready, + recv_x, + recv_x_sf, + recv_weights, + src_ranks, + src_tokens, + src_topk, + l1_fp8, + l1_sf, + l2_fp8, + l2_sf, + l2_x, + l2_x_sf, + combine, + barrier, + out, + ) + if check_capacity: + local_max = recv_counts.max() + dist.all_reduce(local_max, op=dist.ReduceOp.MAX, group=group) + assert local_max.item() <= capacity + return out + + reset_state() + actual = run_pipeline(check_capacity=True) + torch.cuda.synchronize() + dist.barrier(group=group) + + if args.check: + expected = two_kernel.torch_reference( + x_fp8_src, + x_sf_src, + topk_idx_src, + topk_weights_src, + l1_fp8_src, + l1_sf_src, + l2_fp8_src, + l2_sf_src, + group, + activation_clamp, + ) + diff = two_kernel.calc_diff(actual, expected) + assert diff < args.diff_tol, ( + f"rank {local_rank}: diff={diff} exceeds {args.diff_tol}" + ) + print(f"rank {local_rank} check passed, diff={diff:.6f}") + + if args.rep > 0: + reset_state() + for _ in range(args.warmup): + run_pipeline() + reset_state() + latency = do_bench( + run_pipeline, + warmup=0, + rep=args.rep, + post_fn=reset_state, + group=group, + ) + if local_rank == 0: + print( + "tilescale sm90 fp8 mega moe single kernel: " + f"model={model_name} family={shape_family} M={num_tokens} " + f"H={hidden} IH={intermediate_hidden} E={num_experts} " + f"topk={num_topk} capacity={capacity} " + f"epw={config['num_experts_per_wave']} " + f"stages={config['pipeline_stages']} " + f"latency={latency * 1000:.1f} us" + ) + + allocator.close() + dist.destroy_process_group() + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--num-processes", type=int, default=8) + parser.add_argument( + "--model-config", choices=tuple(MODEL_CONFIGS), default="smoke" + ) + parser.add_argument("--hidden", type=int, default=None) + parser.add_argument("--intermediate-hidden", type=int, default=None) + parser.add_argument("--num-experts", type=int, default=None) + parser.add_argument("--num-topk", type=int, default=None) + parser.add_argument("--num-tokens", type=int, default=64) + parser.add_argument("--capacity", type=int, default=None) + parser.add_argument("--experts-per-wave", type=int, default=None) + parser.add_argument("--pipeline-stages", type=int, default=None) + parser.add_argument("--activation-clamp", type=float, default=10.0) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--diff-tol", type=float, default=0.01) + parser.add_argument("--warmup", type=int, default=1) + parser.add_argument("--rep", type=int, default=1) + parser.add_argument("--check", action="store_true") + parser.add_argument("--print-source", action="store_true") + args = parser.parse_args() + torch.multiprocessing.spawn( + main, + args=(args.num_processes, args), + nprocs=args.num_processes, + ) From a86c8c2a9cfe2d58f1f474b6526faa21c8eeca22 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Thu, 27 Aug 2026 14:51:47 +0800 Subject: [PATCH 25/30] Revert the mega MoE scale-pool layout changes Reverts 40688614 (scale-group major `recv_x_sf`) and c8e2f787 (the same treatment for `l2_x_sf`). Both are wrong at M>=2048: `--check` reports a diff of 0.28/0.31 on Flash, in both scatter modes, while M<=1024 passes. 68088edc reports 2.4e-5 at the same size. The two commits were validated at M<=512 only, so the M sweep that followed them measured a kernel that computes wrong results at the top of the range. The 1.2% and 1.1% they claimed are forfeited until the race is understood -- the suspect is the dispatch pull, which now stages the remote scale row in shared memory and scatters it into a transposed pool, where the previous code wrote the pool directly with a single `get_warp`. --- .../mega_moe/example_sm90_fp8_mega_moe.py | 28 ++++++------------- 1 file changed, 9 insertions(+), 19 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 6484b7eef4..d63f198da0 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -234,7 +234,7 @@ def main( m_tasks: T.Tensor((num_experts_per_rank * num_m_blocks,), T.int32), num_m_tasks: T.Tensor((1,), T.int32), recv_x: T.Tensor((num_experts_per_rank, capacity, hidden), T.float8_e4m3fn), - recv_x_sf: T.Tensor((num_experts_per_rank, num_scale_groups, capacity), T.float32), + recv_x_sf: T.Tensor((num_experts_per_rank, capacity, num_scale_groups), T.float32), recv_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), @@ -242,7 +242,7 @@ def main( l1_weight: T.Tensor((num_experts_per_rank, l1_n, hidden), T.float8_e4m3fn), l1_weight_sf: T.Tensor((num_experts_per_rank, l1_n // SCALE_GRANULARITY, hidden // SCALE_GRANULARITY), T.float32), l2_x: T.Tensor((num_experts_per_rank, capacity, l1_n // 2), T.float8_e4m3fn), - l2_x_sf: T.Tensor((num_experts_per_rank, l1_n // (2 * SCALE_GRANULARITY), capacity), T.float32), + l2_x_sf: T.Tensor((num_experts_per_rank, capacity, l1_n // (2 * SCALE_GRANULARITY)), T.float32), barrier: T.Tensor((num_ranks,), T.int32), ): T.annotate_pass_configs({tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) @@ -263,7 +263,6 @@ def main( # access. The producer issues its load before the TMAs so the # latency hides behind TMA issue rather than delaying the arrive. act_sf_shared = T.alloc_shared((pipeline_stages, block_m), T.float32) - pull_sf_shared = T.alloc_shared((dispatch_warps, num_scale_groups), T.float32) stage_barriers = T.alloc_barrier([producer_threads] * pipeline_stages + [num_math_threads] * pipeline_stages) # Routing runs on the math warpgroups, which hold the large budget. @@ -356,12 +355,7 @@ def main( pull_rank = src_ranks[pull_expert, pull_slot] pull_token = src_tokens[pull_expert, pull_slot] T.get_warp(T.address_of(x[pull_token, 0]), T.address_of(recv_x[pull_expert, pull_slot, 0]), hidden, src_pe=pull_rank, unroll_factor=8) - # Keep the remote read contiguous, then scatter into the scale-group - # major pool so the producer can read a block_m column contiguously. - T.get_warp(T.address_of(x_sf[pull_token, 0]), T.address_of(pull_sf_shared[dispatch_warp, 0]), num_scale_groups, src_pe=pull_rank, unroll_factor=8) - T.sync_warp() - for pull_sf_group in T.serial(dispatch_lane, num_scale_groups, warp_size): - recv_x_sf[pull_expert, pull_sf_group, pull_slot] = pull_sf_shared[dispatch_warp, pull_sf_group] + T.get_warp(T.address_of(x_sf[pull_token, 0]), T.address_of(recv_x_sf[pull_expert, pull_slot, 0]), num_scale_groups, src_pe=pull_rank, unroll_factor=8) T.sync_warp() if dispatch_lane == dispatch_leader_lane: T.atom_add(arrivals[pull_expert, pull_slot // block_m], 1, scope="gpu", sem="release") @@ -385,7 +379,7 @@ def main( T.mbarrier_wait_parity(stage_barriers[pipeline_stages + producer_stage], producer_phase ^ 1) producer_sf = T.alloc_local((1,), T.float32) producer_sf[0] = recv_x_sf[ - producer_expert, producer_k, producer_m * block_m + tx - producer_begin] + producer_expert, producer_m * block_m + tx - producer_begin, producer_k] for producer_ks in T.unroll(num_k_sub): producer_sf_k = producer_k * num_k_sub + producer_ks T.tma_copy( @@ -488,8 +482,8 @@ def main( scale[i, scale_group] = T.max(amax[i, scale_group], 1e-4) / FP8_MAX l2_x_sf[ consumer_expert, - consumer_n * num_output_scale_groups + scale_group, consumer_m * block_m + i, + consumer_n * num_output_scale_groups + scale_group, ] = scale[i, scale_group] for i, j in T.Parallel(block_m, block_n // 2): gate[i, j] = T.clamp(gate[i, j] / scale[i, j // SCALE_GRANULARITY], -FP8_MAX, FP8_MAX) @@ -559,7 +553,7 @@ def fused_l2_scatter_reduce_manual_warp_kernel( def main( a: T.Tensor((num_experts_per_rank, capacity, intermediate_hidden), T.float8_e4m3fn), b: T.Tensor((num_experts_per_rank, hidden, intermediate_hidden), T.float8_e4m3fn), - a_sf: T.Tensor((num_experts_per_rank, intermediate_hidden // SCALE_GRANULARITY, capacity), T.float32), + a_sf: T.Tensor((num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), T.float32), b_sf: T.Tensor((num_experts_per_rank, hidden // SCALE_GRANULARITY, intermediate_hidden // SCALE_GRANULARITY), T.float32), recv_counts: T.Tensor((num_experts_per_rank,), T.int32), m_tasks: T.Tensor((num_experts_per_rank * num_m_blocks,), T.int32), @@ -589,7 +583,6 @@ def main( if tx >= producer_begin and tx < producer_end: # WG0 warps 2-3 keep the L2 TMA stages filled. producer_step = T.alloc_var(T.int32, init=0) - producer_sf = T.alloc_local((1,), T.float32) for producer_task in T.serial(bid, num_m_tasks[0] * num_n_blocks, num_sms): producer_n = producer_task % num_n_blocks producer_m_task = m_tasks[producer_task // num_n_blocks] @@ -604,9 +597,6 @@ def main( producer_stage = (producer_step + producer_k) % pipeline_stages producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 T.mbarrier_wait_parity(stage_barriers[pipeline_stages + producer_stage], producer_phase ^ 1) - # Issue the scale read before the TMAs so its latency hides behind TMA - # issue rather than delaying the arrive (measured both ways on L1). - producer_sf[0] = a_sf[producer_expert, producer_k, producer_m * block_m + tx - producer_begin] T.tma_copy( a[ producer_expert, @@ -625,7 +615,7 @@ def main( b_shared[producer_stage, :, :], barrier=stage_barriers[producer_stage], ) - a_sf_shared[producer_stage, tx - producer_begin] = producer_sf[0] + a_sf_shared[producer_stage, tx - producer_begin] = a_sf[producer_expert, producer_m * block_m + tx - producer_begin, producer_k] T.mbarrier_arrive(stage_barriers[producer_stage]) producer_step += num_k_blocks @@ -945,7 +935,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): route_counts = allocator_tensor((num_ranks, num_experts), torch.int32, allocator=allocator) recv_counts = allocator_tensor((num_experts_per_rank,), torch.int32, allocator=allocator) recv_x = allocator_tensor((num_experts_per_rank, capacity, hidden), torch.float8_e4m3fn, allocator=allocator) - recv_x_sf = allocator_tensor((num_experts_per_rank, hidden // SCALE_GRANULARITY, capacity), torch.float32, allocator=allocator) + recv_x_sf = allocator_tensor((num_experts_per_rank, capacity, hidden // SCALE_GRANULARITY), torch.float32, allocator=allocator) recv_weights = allocator_tensor((num_experts_per_rank, capacity), torch.float32, allocator=allocator) src_ranks = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) src_tokens = allocator_tensor((num_experts_per_rank, capacity), torch.int32, allocator=allocator) @@ -955,7 +945,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): m_tasks = allocator_tensor((num_experts_per_rank * ((capacity + 63) // 64),), torch.int32, allocator=allocator) num_m_tasks = allocator_tensor((1,), torch.int32, allocator=allocator) l2_x = allocator_tensor((num_experts_per_rank, capacity, intermediate_hidden), torch.float8_e4m3fn, allocator=allocator) - l2_x_sf = allocator_tensor((num_experts_per_rank, intermediate_hidden // SCALE_GRANULARITY, capacity), torch.float32, allocator=allocator) + l2_x_sf = allocator_tensor((num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), torch.float32, allocator=allocator) combine = allocator_tensor((num_tokens, num_topk, hidden), torch.bfloat16, allocator=allocator) out = allocator_tensor((num_tokens, hidden), torch.bfloat16, allocator=allocator) From ec76e0d1421e68f4573ee5365680914b553756d6 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Thu, 27 Aug 2026 15:04:02 +0800 Subject: [PATCH 26/30] perf(distributed): recalibrate the L2 scatter crossover The `put_warp` scatter was selected from M=256 up, but it only starts winning at M=1024. Measured on Flash (4x H200), direct vs put_warp: M=128 468.2 vs 487.7 M=256 478.8 vs 496.0 M=512 520.1 vs 523.1 M=1024 782.6 vs 763.7 M=2048 1352.5 vs 1216.3 So M=256 and M=512 were paying 3.5% and 0.6% for the wrong branch. The old threshold predates the flat task queue, which changed the per-tile cost balance. Co-Authored-By: Claude Opus 5 (1M context) --- examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index d63f198da0..b7d993099b 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -883,7 +883,9 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): assert requested > 0 and num_experts_per_rank % requested == 0 config["num_experts_per_wave"] = requested l2_scatter = getattr(args, "l2_scatter", "auto") - use_put_warp_scatter = l2_scatter == "warp" or (l2_scatter == "auto" and num_tokens >= 256) + # Crossover measured on Flash (4x H200): direct wins by 4.0%/3.5%/0.6% at + # M=128/256/512, put_warp by 2.4%/10.1% at M=1024/2048. + use_put_warp_scatter = l2_scatter == "warp" or (l2_scatter == "auto" and num_tokens >= 1024) kernel_specs = [ fused_l1_swiglu_manual_warp_kernel( num_tokens, hidden, 2 * intermediate_hidden, num_experts, num_topk, num_ranks, capacity, num_sms, From daa40bc74c749ec13751fbdab6623879ae527a50 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Fri, 28 Aug 2026 11:43:41 +0800 Subject: [PATCH 27/30] feat(distributed): finalize SM90 FP8 MegaMoE example --- examples/distributed/mega_moe/README.md | 92 +++++ .../mega_moe/example_sm90_fp8_mega_moe.py | 363 +++++++----------- .../test_example_sm90_fp8_mega_moe.py | 2 +- 3 files changed, 235 insertions(+), 222 deletions(-) create mode 100644 examples/distributed/mega_moe/README.md diff --git a/examples/distributed/mega_moe/README.md b/examples/distributed/mega_moe/README.md new file mode 100644 index 0000000000..c659d9b79f --- /dev/null +++ b/examples/distributed/mega_moe/README.md @@ -0,0 +1,92 @@ +# SM90 FP8 MegaMoE + +This example implements distributed FP8 MegaMoE with two persistent TileScale +kernels on NVIDIA SM90 GPUs: + +```text +inputs -> dispatch + L1 GEMM + SwiGLU -> L2 GEMM + scatter + reduce -> output +``` + +Let: + +- `M`: tokens per rank; +- `H`: hidden size; +- `I`: intermediate hidden size; +- `E`: global expert count; +- `R`: rank count; +- `K`: experts selected per token; and +- `C`: per-expert capacity. + +Experts are sharded evenly, so each rank owns `E / R` experts. + +## Inputs and Output + +The pipeline receives the following tensors on each rank: + +| Tensor | Shape | Dtype | Description | +| --- | --- | --- | --- | +| `x` | `[M, H]` | FP8 E4M3 | Local input tokens | +| `x_sf` | `[M, H / 128]` | FP32 | Per-128 activation scales | +| `topk_idx` | `[M, K]` | INT32 | Global expert IDs | +| `topk_weights` | `[M, K]` | FP32 | Route weights | +| `l1_weight` | `[E / R, 2I, H]` | FP8 E4M3 | Local gate/up weights | +| `l1_weight_sf` | `[E / R, 2I / 128, H / 128]` | FP32 | L1 per-128 weight scales | +| `l2_weight` | `[E / R, H, I]` | FP8 E4M3 | Local down-projection weights | +| `l2_weight_sf` | `[E / R, H / 128, I / 128]` | FP32 | L2 per-128 weight scales | + +The final output is: + +| Tensor | Shape | Dtype | Description | +| --- | --- | --- | --- | +| `out` | `[M, H]` | BF16 | Sum of the `K` routed expert outputs for each local token | + +## Kernel Boundary + +Kernel 1, `fused_l1_swiglu_manual_warp_kernel`, dispatches tokens to their +expert-owning ranks, computes the gate/up projections and SwiGLU, applies route +weights, and requantizes the intermediate activations. The outputs consumed by kernel 2 are: + +| Tensor | Shape | Dtype | +| --- | --- | --- | +| `l2_x` | `[E / R, C, I]` | FP8 E4M3 | +| `l2_x_sf` | `[E / R, C, I / 128]` | FP32 | +| `recv_counts` | `[E / R]` | INT32 | +| `src_ranks` | `[E / R, C]` | INT32 | +| `src_tokens` | `[E / R, C]` | INT32 | +| `src_topk` | `[E / R, C]` | INT32 | + +Kernel 2, `fused_l2_scatter_reduce_manual_warp_kernel`, consumes these tensors +and the local L2 weights, scatters each routed result back to its source rank, +and reduces the `K` results into `out`. The `combine[M, K, H]` BF16 tensor +is an internal reduction workspace. + +## Run + +The distributed runtime requires peer-accessible SM90 GPUs and a configured +NVIDIA IMEX channel. + +Run a four-GPU correctness smoke test: + +```bash +CUDA_VISIBLE_DEVICES=0,1,2,3 \ + python examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py \ + --num-processes 4 \ + --model-config smoke \ + --num-tokens 32 \ + --capacity 64 \ + --check \ + --rep 0 +``` + +Benchmark the Flash configuration on four GPUs: + +```bash +CUDA_VISIBLE_DEVICES=0,1,2,3 \ + python examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py \ + --num-processes 4 \ + --model-config flash \ + --num-tokens 128 \ + --capacity 64 \ + --warmup 10 \ + --rep 100 +``` diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index b7d993099b..fee7a715a0 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -35,132 +35,6 @@ FP8_MAX = 448.0 SCALE_GRANULARITY = 128 - -def resolve_model_config(args: argparse.Namespace) -> Tuple[str, dict[str, int]]: - model = MODEL_CONFIGS[args.model_config].copy() - overrides = { - "hidden": getattr(args, "hidden", None), - "intermediate_hidden": getattr(args, "intermediate_hidden", None), - "num_experts": getattr(args, "num_experts", None), - "num_topk": getattr(args, "num_topk", None), - } - is_custom = any(value is not None for value in overrides.values()) - model.update({key: value for key, value in overrides.items() if value is not None}) - return ("custom" if is_custom else args.model_config), model - - -def normalize_experts_per_wave(num_experts: int, requested: int) -> int: - requested = min(max(requested, 1), num_experts) - for candidate in range(requested, num_experts + 1): - if num_experts % candidate == 0: - return candidate - return num_experts - - -def select_manual_warp_configs( - hidden: int, - intermediate_hidden: int, - num_tokens: int, - num_topk: int, - num_experts_per_rank: int, - num_sms: int, -) -> Tuple[str, dict[str, int], dict[str, int]]: - """Select the TileScale counterpart of DeepGEMM SM90 schedule families.""" - if 3072 <= hidden < 5120 and 1536 <= intermediate_hidden < 2560: - shape_family = "compact" - elif 5120 <= hidden <= 8192 and 2560 <= intermediate_hidden <= 4096: - shape_family = "wide" - else: - shape_family = "generic" - - routed_tokens = num_tokens * num_topk - high_sm = num_sms >= 100 - - # Measured on Flash (4x H200): three stages ties five at M<=512 and wins - # 0.6%/2.1% at M=2048/8192, so the deeper default is not worth its - # shared memory. - l1_stages = 3 - l2_stages = 3 - generic_experts_per_wave = num_experts_per_rank - if num_experts_per_rank <= routed_tokens <= 4 * num_experts_per_rank: - expected_tokens = (routed_tokens + num_experts_per_rank - 1) // num_experts_per_rank - num_m_blocks = (expected_tokens + 63) // 64 - blocks_per_expert = num_m_blocks * (2 * intermediate_hidden // 256) - requested = min(num_experts_per_rank, (2 * num_sms + blocks_per_expert - 1) // blocks_per_expert) - if blocks_per_expert < num_sms: - max_candidate = min(num_experts_per_rank, 2 * requested) - requested = max( - range(requested, max_candidate + 1), - key=lambda candidate: 1.0 if num_experts_per_rank % candidate == 0 else (num_experts_per_rank % candidate) / candidate, - ) - generic_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, requested) - l1_experts_per_wave = l2_experts_per_wave = generic_experts_per_wave - if high_sm and shape_family == "compact": - if routed_tokens <= 32 * num_experts_per_rank: - l1_stages = l2_stages = 3 - l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 4) - elif 128 * num_experts_per_rank < routed_tokens <= 256 * num_experts_per_rank or routed_tokens > 1024 * num_experts_per_rank: - l1_stages = l2_stages = 4 - l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 32) - elif high_sm and shape_family == "wide": - # BN512/BK256 are profitable in the CUDA kernel, but the manually - # tuned TileScale BN256/BK128 path is faster for the current WGMMA - # lowering and remains the generic Wide schedule. - l1_stages = 4 - if routed_tokens <= 24 * num_experts_per_rank: - # CUDA selects 16 experts here, while TileScale's direct TIR - # scheduler is faster with a shorter four-expert scan on H200. - l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 4) - elif 24 * num_experts_per_rank < routed_tokens <= 48 * num_experts_per_rank: - l1_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 8) - l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 48) - elif routed_tokens > 48 * num_experts_per_rank: - l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 16) - - common = {"block_m": 64, "block_n": 256, "block_k": 128, "threads": 384} - return ( - shape_family, - {**common, "pipeline_stages": l1_stages, "num_experts_per_wave": l1_experts_per_wave}, - {**common, "pipeline_stages": l2_stages, "num_experts_per_wave": l2_experts_per_wave}, - ) - - -def per_token_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: - m, k = x.shape - x_view = x.float().view(m, k // SCALE_GRANULARITY, SCALE_GRANULARITY) - amax = x_view.abs().amax(dim=-1).clamp(1e-4) - scale = amax / FP8_MAX - x_fp8 = (x_view / scale.unsqueeze(-1)).to(torch.float8_e4m3fn) - return x_fp8.view(m, k).contiguous(), scale.contiguous() - - -def block_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: - groups, n, k = x.shape - x_view = x.float().view(groups, n // SCALE_GRANULARITY, SCALE_GRANULARITY, k // SCALE_GRANULARITY, SCALE_GRANULARITY) - amax = x_view.abs().amax(dim=(-1, -3)).clamp(1e-4) - scale = amax / FP8_MAX - x_fp8 = (x_view / scale.unsqueeze(-1).unsqueeze(-3)).to(torch.float8_e4m3fn) - return x_fp8.view(groups, n, k).contiguous(), scale.contiguous() - - -def interleave_gate_up_weights(weight: torch.Tensor, granularity: int = 8) -> torch.Tensor: - groups, n, k = weight.shape - half = n // 2 - gate = weight[:, :half].view(groups, half // granularity, granularity, k) - up = weight[:, half:].view(groups, half // granularity, granularity, k) - return torch.stack((gate, up), dim=2).reshape(groups, n, k).contiguous() - - -def dequantize_per_token(x: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: - m, k = x.shape - return (x.float().view(m, k // SCALE_GRANULARITY, SCALE_GRANULARITY) * scale.unsqueeze(-1)).view(m, k) - - -def dequantize_block(x: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: - groups, n, k = x.shape - x_view = x.float().view(groups, n // SCALE_GRANULARITY, SCALE_GRANULARITY, k // SCALE_GRANULARITY, SCALE_GRANULARITY) - return (x_view * scale.unsqueeze(-1).unsqueeze(-3)).view(groups, n, k) - def fused_l1_swiglu_manual_warp_kernel( num_tokens: int, hidden: int, @@ -177,8 +51,6 @@ def fused_l1_swiglu_manual_warp_kernel( threads: int = 384, pipeline_stages: int = 5, num_experts_per_wave: int | None = None, - frontend_regs_override: int | None = None, - math_regs_override: int | None = None, ): num_experts_per_rank = num_experts // num_ranks num_experts_per_wave = num_experts_per_wave or num_experts_per_rank @@ -211,12 +83,9 @@ def fused_l1_swiglu_manual_warp_kernel( num_l1_scale_groups = l1_n // (2 * SCALE_GRANULARITY) tma_block_n = min(block_n, 256) num_tma_n_blocks = block_n // tma_block_n - # Budgets must leave the CTA register pool some slack: 128*fe + 256*math - # exactly at 65536 (e.g. 32/240) compiles but deadlocks at run time. - # Spilling tracks the frontend budget, not the math one -- 40/48/56 give - # 72/16/0 bytes of spill -- so keep the frontend at 64 for a spill-free build. - frontend_registers = frontend_regs_override or (32 if num_math_threads == 512 else 64) - math_registers = math_regs_override or (112 if num_math_threads == 512 else 192) + # Leave enough register-pool slack for the frontend and math warpgroups. + frontend_registers = 32 if num_math_threads == 512 else 64 + math_registers = 112 if num_math_threads == 512 else 192 dispatch_leader_lane = 0 route_threads = num_math_threads assert route_threads % warp_size == 0 @@ -622,7 +491,7 @@ def main( elif tx >= math_begin: # WG1-2 run WGMMA and scatter their BF16 column pairs remotely. partial = T.alloc_fragment((block_m, block_n), T.float32) - accum = T.alloc_fragment((block_m, block_n), T.bfloat16) + accum = T.alloc_fragment((block_m, block_n), T.float32) act_scale = T.alloc_fragment((block_m,), T.float32) weight_scale = T.alloc_local((2,), T.float32) consumer_step = T.alloc_var(T.int32, init=0) @@ -652,9 +521,7 @@ def main( weight_scale[0] = b_sf[consumer_expert, consumer_n * 2, consumer_k] weight_scale[1] = b_sf[consumer_expert, consumer_n * 2 + 1, consumer_k] for i, j in T.Parallel(block_m, block_n): - accum[i, j] = T.cast(partial[i, j], T.bfloat16) * T.cast( - act_scale[i] * weight_scale[j // 128], T.bfloat16 - ) + accum[i, j] + accum[i, j] = partial[i, j] * act_scale[i] * weight_scale[j // 128] + accum[i, j] T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) consumer_step += num_k_blocks if use_put_warp_scatter: @@ -713,8 +580,8 @@ def main( ): for scatter_col_chunk in T.serial(block_n // math_warpgroups // 8): scatter_col = scatter_wg * (block_n // math_warpgroups) + scatter_col_chunk * 8 + (scatter_lane % 4) * 2 - scatter_value_lo = T.alloc_var(T.uint16, init=T.reinterpret(accum[row, scatter_col], T.uint16)) - scatter_value_hi = T.alloc_var(T.uint16, init=T.reinterpret(accum[row, scatter_col + 1], T.uint16)) + scatter_value_lo = T.alloc_var(T.uint16, init=T.reinterpret(T.cast(accum[row, scatter_col], T.bfloat16), T.uint16)) + scatter_value_hi = T.alloc_var(T.uint16, init=T.reinterpret(T.cast(accum[row, scatter_col + 1], T.bfloat16), T.uint16)) scatter_value = T.alloc_var(T.uint32, init=T.cast(scatter_value_lo, T.uint32) | (T.cast(scatter_value_hi, T.uint32) << 16)) T.st(combine[scatter_dst_token, scatter_dst_topk, consumer_n * block_n + scatter_col], scatter_value, dst_pe=scatter_dst_rank) @@ -804,7 +671,10 @@ def torch_reference( l1_sf_all = _gather_cat(l1_sf, group) l2_all = _gather_cat(l2_fp8, group) l2_sf_all = _gather_cat(l2_sf, group) - x = dequantize_per_token(x_fp8, x_sf) + x_m, x_k = x_fp8.shape + x = ( + x_fp8.float().view(x_m, x_k // SCALE_GRANULARITY, SCALE_GRANULARITY) * x_sf.unsqueeze(-1) + ).view(x_m, x_k) result = torch.zeros((x.size(0), l2_all.size(1)), dtype=torch.float32, device=x.device) for expert_idx in range(l1_all.size(0)): @@ -813,16 +683,39 @@ def torch_reference( continue token_indices = positions[:, 0] topk_slots = positions[:, 1] - l1_weight = dequantize_block(l1_all[expert_idx : expert_idx + 1], l1_sf_all[expert_idx : expert_idx + 1])[0] + l1_weight_fp8 = l1_all[expert_idx : expert_idx + 1] + l1_weight_sf = l1_sf_all[expert_idx : expert_idx + 1] + groups, l1_n, l1_k = l1_weight_fp8.shape + l1_weight_view = l1_weight_fp8.float().view( + groups, l1_n // SCALE_GRANULARITY, SCALE_GRANULARITY, l1_k // SCALE_GRANULARITY, SCALE_GRANULARITY + ) + l1_weight = ( + l1_weight_view * l1_weight_sf.unsqueeze(-1).unsqueeze(-3) + ).view(groups, l1_n, l1_k)[0] gate_up = x[token_indices] @ l1_weight.T gate, up = gate_up.chunk(2, dim=-1) gate = gate.clamp(max=activation_clamp) up = up.clamp(min=-activation_clamp, max=activation_clamp) activated = torch.nn.functional.silu(gate) * up activated *= topk_weights[token_indices, topk_slots].unsqueeze(-1) - activated_fp8, activated_sf = per_token_cast_to_fp8(activated) - activated_dequant = dequantize_per_token(activated_fp8, activated_sf) - l2_weight = dequantize_block(l2_all[expert_idx : expert_idx + 1], l2_sf_all[expert_idx : expert_idx + 1])[0] + activated_m, activated_k = activated.shape + activated_view = activated.float().view( + activated_m, activated_k // SCALE_GRANULARITY, SCALE_GRANULARITY + ) + activated_sf = activated_view.abs().amax(dim=-1).clamp(1e-4) / FP8_MAX + activated_fp8 = (activated_view / activated_sf.unsqueeze(-1)).to(torch.float8_e4m3fn) + activated_dequant = ( + activated_fp8.float() * activated_sf.unsqueeze(-1) + ).view(activated_m, activated_k) + l2_weight_fp8 = l2_all[expert_idx : expert_idx + 1] + l2_weight_sf = l2_sf_all[expert_idx : expert_idx + 1] + groups, l2_n, l2_k = l2_weight_fp8.shape + l2_weight_view = l2_weight_fp8.float().view( + groups, l2_n // SCALE_GRANULARITY, SCALE_GRANULARITY, l2_k // SCALE_GRANULARITY, SCALE_GRANULARITY + ) + l2_weight = ( + l2_weight_view * l2_weight_sf.unsqueeze(-1).unsqueeze(-3) + ).view(groups, l2_n, l2_k)[0] contribution = (activated_dequant @ l2_weight.T).to(torch.bfloat16).float() result.index_add_(0, token_indices, contribution) @@ -841,6 +734,101 @@ def allocator_tensor(shape, dtype, allocator): def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): + def per_token_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: + m, k = x.shape + x_view = x.float().view(m, k // SCALE_GRANULARITY, SCALE_GRANULARITY) + scale = x_view.abs().amax(dim=-1).clamp(1e-4) / FP8_MAX + return (x_view / scale.unsqueeze(-1)).to(torch.float8_e4m3fn).view(m, k).contiguous(), scale.contiguous() + + def block_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: + groups, n, k = x.shape + x_view = x.float().view(groups, n // SCALE_GRANULARITY, SCALE_GRANULARITY, k // SCALE_GRANULARITY, SCALE_GRANULARITY) + scale = x_view.abs().amax(dim=(-1, -3)).clamp(1e-4) / FP8_MAX + return (x_view / scale.unsqueeze(-1).unsqueeze(-3)).to(torch.float8_e4m3fn).view(groups, n, k).contiguous(), scale.contiguous() + + def resolve_model_config(args: argparse.Namespace) -> Tuple[str, dict[str, int]]: + model = MODEL_CONFIGS[args.model_config].copy() + overrides = { + "hidden": getattr(args, "hidden", None), + "intermediate_hidden": getattr(args, "intermediate_hidden", None), + "num_experts": getattr(args, "num_experts", None), + "num_topk": getattr(args, "num_topk", None), + } + is_custom = any(value is not None for value in overrides.values()) + model.update({key: value for key, value in overrides.items() if value is not None}) + return ("custom" if is_custom else args.model_config), model + + def select_manual_warp_configs( + hidden: int, + intermediate_hidden: int, + num_tokens: int, + num_topk: int, + num_experts_per_rank: int, + num_sms: int, + ) -> Tuple[str, dict[str, int], dict[str, int]]: + """Select the TileScale counterpart of DeepGEMM SM90 schedule families.""" + if 3072 <= hidden < 5120 and 1536 <= intermediate_hidden < 2560: + shape_family = "compact" + elif 5120 <= hidden <= 8192 and 2560 <= intermediate_hidden <= 4096: + shape_family = "wide" + else: + shape_family = "generic" + + routed_tokens = num_tokens * num_topk + high_sm = num_sms >= 100 + + # Three stages balance pipeline depth and shared-memory use for the default path. + l1_stages = 3 + l2_stages = 3 + generic_experts_per_wave = num_experts_per_rank + def normalize_experts_per_wave(num_experts: int, requested: int) -> int: + requested = min(max(requested, 1), num_experts) + for candidate in range(requested, num_experts + 1): + if num_experts % candidate == 0: + return candidate + return num_experts + if num_experts_per_rank <= routed_tokens <= 4 * num_experts_per_rank: + expected_tokens = (routed_tokens + num_experts_per_rank - 1) // num_experts_per_rank + num_m_blocks = (expected_tokens + 63) // 64 + blocks_per_expert = num_m_blocks * (2 * intermediate_hidden // 256) + requested = min(num_experts_per_rank, (2 * num_sms + blocks_per_expert - 1) // blocks_per_expert) + if blocks_per_expert < num_sms: + max_candidate = min(num_experts_per_rank, 2 * requested) + requested = max( + range(requested, max_candidate + 1), + key=lambda candidate: 1.0 if num_experts_per_rank % candidate == 0 else (num_experts_per_rank % candidate) / candidate, + ) + generic_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, requested) + l1_experts_per_wave = l2_experts_per_wave = generic_experts_per_wave + if high_sm and shape_family == "compact": + if routed_tokens <= 32 * num_experts_per_rank: + l1_stages = l2_stages = 3 + l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 4) + elif 128 * num_experts_per_rank < routed_tokens <= 256 * num_experts_per_rank or routed_tokens > 1024 * num_experts_per_rank: + l1_stages = l2_stages = 4 + l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 32) + elif high_sm and shape_family == "wide": + # BN512/BK256 are profitable in the CUDA kernel, but the manually + # tuned TileScale BN256/BK128 path is faster for the current WGMMA + # lowering and remains the generic Wide schedule. + l1_stages = 4 + if routed_tokens <= 24 * num_experts_per_rank: + # CUDA selects 16 experts here, while TileScale's direct TIR + # scheduler is faster with a shorter four-expert scan on H200. + l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 4) + elif 24 * num_experts_per_rank < routed_tokens <= 48 * num_experts_per_rank: + l1_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 8) + l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 48) + elif routed_tokens > 48 * num_experts_per_rank: + l1_experts_per_wave = l2_experts_per_wave = normalize_experts_per_wave(num_experts_per_rank, 16) + + common = {"block_m": 64, "block_n": 256, "block_k": 128, "threads": 384} + return ( + shape_family, + {**common, "pipeline_stages": l1_stages, "num_experts_per_wave": l1_experts_per_wave}, + {**common, "pipeline_stages": l2_stages, "num_experts_per_wave": l2_experts_per_wave}, + ) + model_name, model = resolve_model_config(args) hidden = model["hidden"] intermediate_hidden = model["intermediate_hidden"] @@ -861,7 +849,7 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): rank, num_ranks, group = init_dist(local_rank, num_local_ranks) assert rank == local_rank and num_ranks == num_local_ranks - num_sms = args.num_sms or torch.cuda.get_device_properties(local_rank).multi_processor_count + num_sms = torch.cuda.get_device_properties(local_rank).multi_processor_count allocator = get_allocator( size=_allocator_size_bytes(num_tokens, hidden, intermediate_hidden, num_experts_per_rank, num_topk, capacity), device=f"cuda:{local_rank}", @@ -873,25 +861,12 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): ) shape_family, l1_config, l2_config = select_manual_warp_configs(hidden, intermediate_hidden, num_tokens, num_topk, num_experts_per_rank, num_sms) - if args.l1_block_k is not None: - l1_config["block_k"] = args.l1_block_k - if args.l1_stages is not None: - l1_config["pipeline_stages"] = args.l1_stages - for phase, config in (("l1", l1_config), ("l2", l2_config)): - requested = getattr(args, f"{phase}_experts_per_wave", None) - if requested is not None: - assert requested > 0 and num_experts_per_rank % requested == 0 - config["num_experts_per_wave"] = requested - l2_scatter = getattr(args, "l2_scatter", "auto") - # Crossover measured on Flash (4x H200): direct wins by 4.0%/3.5%/0.6% at - # M=128/256/512, put_warp by 2.4%/10.1% at M=1024/2048. - use_put_warp_scatter = l2_scatter == "warp" or (l2_scatter == "auto" and num_tokens >= 1024) + # Packed warp scatter amortizes its shared-memory staging for larger token batches. + use_put_warp_scatter = num_tokens >= 1024 kernel_specs = [ fused_l1_swiglu_manual_warp_kernel( num_tokens, hidden, 2 * intermediate_hidden, num_experts, num_topk, num_ranks, capacity, num_sms, - activation_clamp=activation_clamp, - frontend_regs_override=args.l1_frontend_regs, - math_regs_override=args.l1_math_regs, **l1_config, + activation_clamp=activation_clamp, **l1_config, ), fused_l2_scatter_reduce_manual_warp_kernel( num_tokens, hidden, intermediate_hidden, num_experts_per_rank, num_topk, num_ranks, capacity, num_sms, **l2_config, @@ -927,7 +902,11 @@ def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): x_sf = allocator_tensor(x_sf_src.shape, x_sf_src.dtype, allocator=allocator).copy_(x_sf_src) topk_idx = allocator_tensor(topk_idx_src.shape, topk_idx_src.dtype, allocator=allocator).copy_(topk_idx_src) topk_weights = allocator_tensor(topk_weights_src.shape, topk_weights_src.dtype, allocator=allocator).copy_(topk_weights_src) - l1_fp8_kernel = interleave_gate_up_weights(l1_fp8_src) + groups, l1_n, l1_k = l1_fp8_src.shape + half = l1_n // 2 + gate = l1_fp8_src[:, :half].view(groups, half // 8, 8, l1_k) + up = l1_fp8_src[:, half:].view(groups, half // 8, 8, l1_k) + l1_fp8_kernel = torch.stack((gate, up), dim=2).reshape(groups, l1_n, l1_k).contiguous() l1_fp8 = allocator_tensor(l1_fp8_kernel.shape, l1_fp8_kernel.dtype, allocator=allocator).copy_(l1_fp8_kernel) del l1_fp8_kernel l1_sf = allocator_tensor(l1_sf_src.shape, l1_sf_src.dtype, allocator=allocator).copy_(l1_sf_src) @@ -1011,55 +990,6 @@ def run_pipeline(check_capacity: bool = False): f"latency={latency * 1000:.1f} us" ) - if args.profile_phases > 0: - reset_state() - for _ in range(args.warmup): - run_pipeline() - reset_state() - - samples = [] - for _ in range(args.profile_phases): - dist.barrier(group=group) - events = [torch.cuda.Event(enable_timing=True) for _ in range(3)] - events[0].record() - fused_l1( - x, x_sf, topk_idx, topk_weights, route_counts, recv_counts, route_slots, arrivals, - m_tasks, num_m_tasks, recv_x, recv_x_sf, recv_weights, src_ranks, src_tokens, src_topk, - l1_fp8, l1_sf, l2_x, l2_x_sf, barrier, - ) - events[1].record() - fused_l2( - l2_x, l2_fp8, l2_x_sf, l2_sf, recv_counts, m_tasks, num_m_tasks, src_ranks, - src_tokens, src_topk, combine, barrier, out, - ) - events[2].record() - events[2].synchronize() - local = torch.tensor( - [events[0].elapsed_time(events[1]), events[1].elapsed_time(events[2])], - dtype=torch.float32, - device="cuda", - ) - gathered = [torch.empty_like(local) for _ in range(num_ranks)] - dist.all_gather(gathered, local, group=group) - if local_rank == 0: - samples.append(torch.stack(gathered).cpu()) - reset_state() - - if local_rank == 0: - stacked = torch.stack(samples) - max_rank_median = stacked.max(dim=1).values.median(dim=0).values * 1000 - rank_medians = stacked.median(dim=0).values * 1000 - print( - f"phase profile: samples={args.profile_phases} max-rank median " - f"l1={max_rank_median[0]:.1f} us l2={max_rank_median[1]:.1f} us " - f"total={max_rank_median[0] + max_rank_median[1]:.1f} us" - ) - for phase_idx, phase_name in enumerate(("l1", "l2")): - rank_values = ", ".join( - f"r{r}={rank_medians[r, phase_idx]:.1f}" for r in range(num_ranks) - ) - print(f"phase profile {phase_name} rank medians (us): {rank_values}") - allocator.close() dist.destroy_process_group() @@ -1074,20 +1004,11 @@ def run_pipeline(check_capacity: bool = False): parser.add_argument("--num-topk", type=int, default=None) parser.add_argument("--num-tokens", type=int, default=64) parser.add_argument("--capacity", type=int, default=None) - parser.add_argument("--l1-experts-per-wave", type=int, default=None) - parser.add_argument("--l2-experts-per-wave", type=int, default=None) - parser.add_argument("--l2-scatter", choices=("auto", "direct", "warp"), default="auto") parser.add_argument("--activation-clamp", type=float, default=10.0) parser.add_argument("--seed", type=int, default=0) parser.add_argument("--diff-tol", type=float, default=0.01) parser.add_argument("--warmup", type=int, default=1) parser.add_argument("--rep", type=int, default=1) - parser.add_argument("--profile-phases", type=int, default=0) - parser.add_argument("--num-sms", type=int, default=None) - parser.add_argument("--l1-block-k", type=int, default=None) - parser.add_argument("--l1-frontend-regs", type=int, default=None) - parser.add_argument("--l1-math-regs", type=int, default=None) - parser.add_argument("--l1-stages", type=int, default=None) parser.add_argument("--check", action="store_true") parser.add_argument("--print-source", action="store_true") args = parser.parse_args() diff --git a/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py index 02b7cbbc8b..d0fdeba4ee 100644 --- a/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/test_example_sm90_fp8_mega_moe.py @@ -39,7 +39,7 @@ def test_custom_model_config_and_schedule(): "block_n": 256, "block_k": 128, "threads": 384, - "pipeline_stages": 5, + "pipeline_stages": 3, "num_experts_per_wave": 16, } assert l2 == {**l1, "pipeline_stages": 3} From 3f0eaa0937bb3871b2314fbdb6fb678134b5b421 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Fri, 28 Aug 2026 13:47:25 +0800 Subject: [PATCH 28/30] refactor(distributed): localize MegaMoE helpers --- .../mega_moe/example_sm90_fp8_mega_moe.py | 29 +++++++++---------- 1 file changed, 14 insertions(+), 15 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index fee7a715a0..98efbe7fb1 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -649,11 +649,6 @@ def _allocator_size_bytes( return (total_bytes + 2**20 - 1) // 2**20 * 2**20 -def _gather_cat(tensor: torch.Tensor, group: dist.ProcessGroup) -> torch.Tensor: - gathered = [torch.empty_like(tensor) for _ in range(dist.get_world_size(group))] - dist.all_gather(gathered, tensor, group=group) - return torch.cat(gathered, dim=0) - def torch_reference( x_fp8: torch.Tensor, @@ -667,6 +662,10 @@ def torch_reference( group: dist.ProcessGroup, activation_clamp: float, ) -> torch.Tensor: + def _gather_cat(tensor: torch.Tensor, group: dist.ProcessGroup) -> torch.Tensor: + gathered = [torch.empty_like(tensor) for _ in range(dist.get_world_size(group))] + dist.all_gather(gathered, tensor, group=group) + return torch.cat(gathered, dim=0) l1_all = _gather_cat(l1_fp8, group) l1_sf_all = _gather_cat(l1_sf, group) l2_all = _gather_cat(l2_fp8, group) @@ -722,16 +721,6 @@ def torch_reference( return result.to(torch.bfloat16) -def calc_diff(x: torch.Tensor, y: torch.Tensor) -> float: - x, y = x.double(), y.double() - return (1 - 2 * (x * y).sum() / (x.square() + y.square()).sum()).item() - - -def allocator_tensor(shape, dtype, allocator): - if dtype == torch.float8_e4m3fn: - return tilelang.tensor(shape, torch.uint8, allocator=allocator).view(dtype) - return tilelang.tensor(shape, dtype, allocator=allocator) - def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): def per_token_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: @@ -897,6 +886,11 @@ def normalize_experts_per_wave(num_experts: int, requested: int) -> int: # barrier_blocks currently lowers its byte offset as int32, so keep this # allocation before the multi-gigabyte Pro-model weights. + def allocator_tensor(shape, dtype, allocator): + if dtype == torch.float8_e4m3fn: + return tilelang.tensor(shape, torch.uint8, allocator=allocator).view(dtype) + return tilelang.tensor(shape, dtype, allocator=allocator) + barrier = allocator_tensor((num_ranks,), torch.int32, allocator=allocator) x = allocator_tensor(x_fp8_src.shape, x_fp8_src.dtype, allocator=allocator).copy_(x_fp8_src) x_sf = allocator_tensor(x_sf_src.shape, x_sf_src.dtype, allocator=allocator).copy_(x_sf_src) @@ -965,6 +959,11 @@ def run_pipeline(check_capacity: bool = False): dist.barrier(group=group) if args.check: + + def calc_diff(x: torch.Tensor, y: torch.Tensor) -> float: + x, y = x.double(), y.double() + return (1 - 2 * (x * y).sum() / (x.square() + y.square()).sum()).item() + expected = torch_reference( x_fp8_src, x_sf_src, topk_idx_src, topk_weights_src, l1_fp8_src, l1_sf_src, l2_fp8_src, l2_sf_src, group, activation_clamp, From 2eb007a59103ef62c2a472895c4bff7a1f901388 Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Sat, 29 Aug 2026 00:32:57 +0800 Subject: [PATCH 29/30] refactor(distributed): remove single-kernel MegaMoE prototype --- ...example_sm90_fp8_mega_moe_single_kernel.py | 1327 ----------------- 1 file changed, 1327 deletions(-) delete mode 100644 examples/distributed/mega_moe/example_sm90_fp8_mega_moe_single_kernel.py diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe_single_kernel.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe_single_kernel.py deleted file mode 100644 index c46096c6fb..0000000000 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe_single_kernel.py +++ /dev/null @@ -1,1327 +0,0 @@ -"""Experimental single-kernel SM90 FP8 Mega MoE using TileScale. - -Dedicated dispatch, GEMM, and combine warps form a cross-wave pipeline. L1 -epilogues publish per-(expert, M-block) readiness for L2, while rank-wave -completion flags let combine consume wave N as the GEMM warps enter wave N+1. -The established two-kernel example remains the stable baseline. -""" - -from __future__ import annotations - -import argparse -import os -from typing import Tuple - -import torch -import torch.distributed as dist -import torch.multiprocessing - -import tilelang -import tilelang.language as T -from tilelang.distributed.allocator import get_allocator -from tilelang.distributed.bench import do_bench -from tilelang.distributed.host import init_dist - -import example_sm90_fp8_mega_moe as two_kernel - - -os.environ.setdefault("NCCL_DEBUG", "ERROR") - -MODEL_CONFIGS = two_kernel.MODEL_CONFIGS -FP8_MAX = two_kernel.FP8_MAX -SCALE_GRANULARITY = two_kernel.SCALE_GRANULARITY - - -def select_single_kernel_config( - hidden: int, - intermediate_hidden: int, - num_tokens: int, - num_topk: int, - num_experts_per_rank: int, - num_sms: int, -) -> Tuple[str, dict[str, int]]: - """Collapse the validated L1/L2 schedules into one compatible schedule.""" - family, l1, l2 = two_kernel.select_manual_warp_configs( - hidden, - intermediate_hidden, - num_tokens, - num_topk, - num_experts_per_rank, - num_sms, - ) - for key in ("block_m", "block_n", "block_k", "threads"): - assert l1[key] == l2[key] - preferred_wave_size = 16 if family == "compact" else 24 - preferred_wave_size = min(preferred_wave_size, num_experts_per_rank) - wave_size = next( - size - for size in range(preferred_wave_size, num_experts_per_rank + 1) - if num_experts_per_rank % size == 0 - ) - return family, { - "block_m": l1["block_m"], - "block_n": l1["block_n"], - "block_k": l1["block_k"], - "threads": 512, - # The fused allocation includes both epilogues. Four stages keeps its - # shared-memory footprint below the SM90 per-CTA limit. - "pipeline_stages": min(max(l1["pipeline_stages"], l2["pipeline_stages"]), 4), - "num_experts_per_wave": wave_size, - } - - -def fused_single_kernel( - num_tokens: int, - hidden: int, - intermediate_hidden: int, - num_experts: int, - num_topk: int, - num_ranks: int, - capacity: int, - num_sms: int, - activation_clamp: float = 10.0, - block_m: int = 64, - block_n: int = 256, - block_k: int = 128, - threads: int = 512, - pipeline_stages: int = 3, - num_experts_per_wave: int | None = None, -): - num_experts_per_rank = num_experts // num_ranks - num_experts_per_wave = num_experts_per_wave or num_experts_per_rank - assert num_experts_per_rank % num_experts_per_wave == 0 - assert block_m == 64 and block_n == 256 and block_k == 128 - assert threads == 512 and pipeline_stages <= 4 - assert hidden % block_n == 0 and intermediate_hidden % block_k == 0 - - l1_n = 2 * intermediate_hidden - num_scale_groups = hidden // SCALE_GRANULARITY - num_routes = num_tokens * num_topk - num_m_blocks = T.ceildiv(capacity, block_m) - l1_num_n_blocks = l1_n // block_n - l2_num_n_blocks = hidden // block_n - l1_num_k_blocks = hidden // block_k - l2_num_k_blocks = intermediate_hidden // block_k - num_expert_waves = num_experts_per_rank // num_experts_per_wave - # Task cursors flatten (expert, M block, N block) so L1 and L2 can advance - # independently once their block-level readiness checks are enabled. - l1_total_tasks = num_experts_per_rank * num_m_blocks * l1_num_n_blocks - l2_total_tasks = num_experts_per_rank * num_m_blocks * l2_num_n_blocks - l1_total_rounds = T.ceildiv(l1_total_tasks, num_sms) - l2_total_rounds = T.ceildiv(l2_total_tasks, num_sms) - - num_output_scale_groups = block_n // (2 * SCALE_GRANULARITY) - num_l1_scale_groups = l1_n // (2 * SCALE_GRANULARITY) - combine_block_n = 128 - num_combine_n_blocks = hidden // combine_block_n - # Eight chunks amortize each wave wait while keeping enough tasks to fill - # the four combine warps on all SMs. - combine_n_blocks_per_task = 8 - num_combine_groups = T.ceildiv( - num_combine_n_blocks, combine_n_blocks_per_task - ) - combine_values_per_lane = combine_block_n // 32 - - warp_size = 32 - dispatch_threads = 64 - producer_begin = dispatch_threads - producer_threads = 64 - producer_end = producer_begin + producer_threads - math_begin = producer_end - num_math_threads = 256 - math_end = math_begin + num_math_threads - # Keep the validated producer/math warp IDs unchanged and append combine. - combine_begin = math_end - combine_threads = 128 - combine_end = combine_begin + combine_threads - num_combine_warps = combine_threads // warp_size - math_warps = num_math_threads // warp_size - rows_per_math_warp = block_m // math_warps - route_threads = 256 - - @T.prim_func - def main( - x: T.Tensor((num_tokens, hidden), T.float8_e4m3fn), - x_sf: T.Tensor((num_tokens, num_scale_groups), T.float32), - topk_idx: T.Tensor((num_tokens, num_topk), T.int32), - topk_weights: T.Tensor((num_tokens, num_topk), T.float32), - route_counts: T.Tensor((num_ranks, num_experts), T.int32), - recv_counts: T.Tensor((num_experts_per_rank,), T.int32), - route_slots: T.Tensor((num_tokens, num_topk), T.int32), - dispatch_arrivals: T.Tensor( - (num_experts_per_rank, T.ceildiv(capacity, block_m)), T.uint32 - ), - l2_arrivals: T.Tensor( - (num_experts_per_rank, T.ceildiv(capacity, block_m)), T.uint32 - ), - l2_task_ready: T.Tensor( - (num_experts_per_rank, T.ceildiv(capacity, block_m), l2_num_n_blocks), T.uint32 - ), - recv_x: T.Tensor( - (num_experts_per_rank, capacity, hidden), T.float8_e4m3fn - ), - recv_x_sf: T.Tensor( - (num_experts_per_rank, capacity, num_scale_groups), T.float32 - ), - recv_weights: T.Tensor((num_experts_per_rank, capacity), T.float32), - src_ranks: T.Tensor((num_experts_per_rank, capacity), T.int32), - src_tokens: T.Tensor((num_experts_per_rank, capacity), T.int32), - src_topk: T.Tensor((num_experts_per_rank, capacity), T.int32), - l1_weight: T.Tensor( - (num_experts_per_rank, 2 * intermediate_hidden, hidden), - T.float8_e4m3fn, - ), - l1_weight_sf: T.Tensor( - ( - num_experts_per_rank, - 2 * intermediate_hidden // SCALE_GRANULARITY, - hidden // SCALE_GRANULARITY, - ), - T.float32, - ), - l2_weight: T.Tensor( - (num_experts_per_rank, hidden, intermediate_hidden), - T.float8_e4m3fn, - ), - l2_weight_sf: T.Tensor( - ( - num_experts_per_rank, - hidden // SCALE_GRANULARITY, - intermediate_hidden // SCALE_GRANULARITY, - ), - T.float32, - ), - l2_x: T.Tensor( - (num_experts_per_rank, capacity, intermediate_hidden), - T.float8_e4m3fn, - ), - l2_x_sf: T.Tensor( - ( - num_experts_per_rank, - capacity, - intermediate_hidden // SCALE_GRANULARITY, - ), - T.float32, - ), - combine: T.Tensor((num_tokens, num_topk, hidden), T.bfloat16), - barrier: T.Tensor((num_ranks,), T.int32), - out: T.Tensor((num_tokens, hidden), T.bfloat16), - ): - T.annotate_pass_configs({tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) - with T.Kernel(num_sms, threads=threads) as bid: - tx = T.get_thread_binding() - src_rank = T.alloc_local((1,), T.int32) - src_rank[0] = T.get_rank() - - a_shared = T.alloc_shared( - (pipeline_stages, block_m, block_k), T.float8_e4m3fn - ) - b_shared = T.alloc_shared( - (pipeline_stages, block_n, block_k), T.float8_e4m3fn - ) - l2_a_sf_shared = T.alloc_shared( - (pipeline_stages, block_m), T.float32 - ) - l1_out_shared = T.alloc_shared( - (block_m, block_n // 2), T.float8_e4m3fn - ) - l2_out_shared = T.alloc_shared((block_m, block_n), T.bfloat16) - stage_barriers = T.alloc_barrier( - [producer_threads] * pipeline_stages - + [num_math_threads] * pipeline_stages - ) - - if bid == 0: - if tx < route_threads: - for reset_wave in T.serial(T.ceildiv(num_experts, route_threads)): - reset_expert = tx + reset_wave * route_threads - if reset_expert < num_experts: - route_counts[src_rank[0], reset_expert] = 0 - T.sync_threads(7, route_threads) - - for assign_wave in T.serial(T.ceildiv(num_routes, route_threads)): - assign_route = tx + assign_wave * route_threads - if assign_route < num_routes: - assign_token = assign_route // num_topk - assign_topk = assign_route % num_topk - assign_expert = topk_idx[assign_token, assign_topk] - if assign_expert >= 0 and assign_expert < num_experts: - route_slots[assign_token, assign_topk] = T.atomic_add( - route_counts[src_rank[0], assign_expert], - 1, - memory_order="relaxed", - return_prev=True, - ) - else: - route_slots[assign_token, assign_topk] = -1 - T.sync_threads(7, route_threads) - - for publish_wave in T.serial( - T.ceildiv(num_experts * num_ranks, route_threads) - ): - publish_idx = tx + publish_wave * route_threads - if publish_idx < num_experts * num_ranks: - publish_rank = publish_idx // num_experts - publish_expert = publish_idx % num_experts - if publish_rank != src_rank[0]: - T.st( - route_counts[src_rank[0], publish_expert], - route_counts[src_rank[0], publish_expert], - dst_pe=publish_rank, - ) - - T.barrier_blocks(barrier[0]) - - if tx < route_threads: - for count_wave in T.serial( - T.ceildiv(num_experts_per_rank, route_threads) - ): - local_expert = tx + count_wave * route_threads - if local_expert < num_experts_per_rank: - recv_count = T.alloc_var(T.int32, init=0) - recv_expert = ( - src_rank[0] * num_experts_per_rank + local_expert - ) - for count_rank in T.serial(num_ranks): - recv_count += route_counts[count_rank, recv_expert] - recv_counts[local_expert] = recv_count - - for prefix_wave in T.serial(T.ceildiv(num_routes, route_threads)): - prefix_route = tx + prefix_wave * route_threads - if prefix_route < num_routes: - prefix_token = prefix_route // num_topk - prefix_topk = prefix_route % num_topk - prefix_expert = topk_idx[prefix_token, prefix_topk] - prefix_slot = T.alloc_var( - T.int32, - init=route_slots[prefix_token, prefix_topk], - ) - if ( - prefix_expert >= 0 - and prefix_expert < num_experts - and prefix_slot >= 0 - ): - for prefix_rank in T.serial(num_ranks): - if prefix_rank < src_rank[0]: - prefix_slot += route_counts[ - prefix_rank, prefix_expert - ] - route_slots[prefix_token, prefix_topk] = prefix_slot - - T.sync_grid() - - if tx < math_begin or tx >= math_end: - T.dec_max_nreg(48) - else: - T.inc_max_nreg(208) - - dispatch_warp = tx // warp_size - dispatch_lane = tx % warp_size - if tx < dispatch_threads: - for metadata_wave in T.serial( - T.ceildiv(num_routes, num_sms * dispatch_threads) - ): - metadata_route = ( - bid * dispatch_threads - + tx - + metadata_wave * num_sms * dispatch_threads - ) - if metadata_route < num_routes: - metadata_token = metadata_route // num_topk - metadata_topk = metadata_route % num_topk - metadata_expert = topk_idx[metadata_token, metadata_topk] - metadata_slot = route_slots[metadata_token, metadata_topk] - if ( - metadata_expert >= 0 - and metadata_slot >= 0 - and metadata_slot < capacity - ): - metadata_rank = ( - metadata_expert // num_experts_per_rank - ) - metadata_local_expert = ( - metadata_expert % num_experts_per_rank - ) - T.st( - recv_weights[ - metadata_local_expert, metadata_slot - ], - topk_weights[metadata_token, metadata_topk], - dst_pe=metadata_rank, - ) - T.st( - src_tokens[metadata_local_expert, metadata_slot], - metadata_token, - dst_pe=metadata_rank, - ) - T.st( - src_topk[metadata_local_expert, metadata_slot], - metadata_topk, - dst_pe=metadata_rank, - ) - T.st( - src_ranks[metadata_local_expert, metadata_slot], - src_rank[0], - scope="sys", - sem="release", - dst_pe=metadata_rank, - ) - - for pull_wave in T.serial( - T.ceildiv( - num_experts_per_rank * capacity, - num_sms * (dispatch_threads // warp_size), - ) - ): - pull_idx = ( - bid * (dispatch_threads // warp_size) - + dispatch_warp - + pull_wave - * num_sms - * (dispatch_threads // warp_size) - ) - pull_expert = pull_idx // capacity - pull_slot = pull_idx % capacity - if ( - pull_expert < num_experts_per_rank - and pull_slot < recv_counts[pull_expert] - ): - if dispatch_lane == 0: - T.wait_ge( - src_ranks[pull_expert, pull_slot], - 0, - scope=T.WaitScope.SYS, - semantics=T.WaitSemantics.ACQUIRE, - ) - T.sync_warp() - pull_rank = src_ranks[pull_expert, pull_slot] - pull_token = src_tokens[pull_expert, pull_slot] - T.get_warp( - T.address_of(x[pull_token, 0]), - T.address_of(recv_x[pull_expert, pull_slot, 0]), - hidden, - src_pe=pull_rank, - unroll_factor=8, - ) - T.get_warp( - T.address_of(x_sf[pull_token, 0]), - T.address_of(recv_x_sf[pull_expert, pull_slot, 0]), - num_scale_groups, - src_pe=pull_rank, - unroll_factor=8, - ) - T.sync_warp() - if dispatch_lane == 0: - T.atom_add( - dispatch_arrivals[ - pull_expert, pull_slot // block_m - ], - 1, - scope="gpu", - sem="release", - ) - - if tx >= combine_begin and tx < combine_end: - combine_warp = (tx - combine_begin) // warp_size - combine_lane = (tx - combine_begin) % warp_size - combine_accum = T.alloc_local( - ( - combine_n_blocks_per_task, - combine_values_per_lane, - ), - T.float32, - ) - num_combine_tasks = num_tokens * num_combine_groups - for combine_round in T.serial( - T.ceildiv( - num_combine_tasks, - num_sms * num_combine_warps, - ) - ): - combine_task = ( - bid * num_combine_warps - + combine_warp - + combine_round * num_sms * num_combine_warps - ) - if combine_task < num_combine_tasks: - combine_token = combine_task // num_combine_groups - combine_group = combine_task % num_combine_groups - for group_block in T.unroll( - combine_n_blocks_per_task - ): - for value_idx in T.unroll( - combine_values_per_lane - ): - combine_accum[group_block, value_idx] = 0.0 - - for expert_wave in T.serial(num_expert_waves): - wave_begin = expert_wave * num_experts_per_wave - wave_end = wave_begin + num_experts_per_wave - for topk_slot in T.serial(num_topk): - route_expert = topk_idx[ - combine_token, topk_slot - ] - route_slot = route_slots[ - combine_token, topk_slot - ] - route_local_expert = ( - route_expert % num_experts_per_rank - ) - route_rank = ( - route_expert // num_experts_per_rank - ) - if ( - route_expert >= 0 - and route_slot >= 0 - and route_slot < capacity - and route_local_expert >= wave_begin - and route_local_expert < wave_end - ): - for group_block in T.unroll( - combine_n_blocks_per_task - ): - combine_n_block = ( - combine_group - * combine_n_blocks_per_task - + group_block - ) - if combine_n_block < num_combine_n_blocks: - if combine_lane == 0: - T.wait_ge( - l2_task_ready[ - route_local_expert, - route_slot // block_m, - combine_n_block - // (block_n // combine_block_n), - ], - 1, - scope=T.WaitScope.SYS, - semantics=T.WaitSemantics.ACQUIRE, - ) - T.sync_warp() - for value_idx in T.unroll( - combine_values_per_lane - ): - combine_col = ( - combine_n_block - * combine_block_n - + value_idx * warp_size - + combine_lane - ) - combine_accum[ - group_block, value_idx - ] += T.cast( - combine[ - combine_token, - topk_slot, - combine_col, - ], - T.float32, - ) - - for group_block in T.unroll( - combine_n_blocks_per_task - ): - combine_n_block = ( - combine_group * combine_n_blocks_per_task - + group_block - ) - if combine_n_block < num_combine_n_blocks: - for value_idx in T.unroll( - combine_values_per_lane - ): - combine_col = ( - combine_n_block * combine_block_n - + value_idx * warp_size - + combine_lane - ) - out[combine_token, combine_col] = T.cast( - combine_accum[ - group_block, value_idx - ], - T.bfloat16, - ) - - if tx >= producer_begin and tx < producer_end: - producer_step = T.alloc_var(T.int32, init=0) - l1_task = T.alloc_var(T.int32, init=bid) - l2_task = T.alloc_var(T.int32, init=bid) - - # SM90 keeps one producer warpgroup and advances each phase's - # flattened task cursor independently. The L2 cursor waits on - # the per-expert/M-block L1 readiness counter below. - for l1_round in T.serial(l1_total_rounds): - if l1_task < l1_total_tasks: - l1_task_offset = T.alloc_var(T.int32, init=l1_task) - l1_expert = l1_task_offset // (num_m_blocks * l1_num_n_blocks) - l1_task_offset -= l1_expert * num_m_blocks * l1_num_n_blocks - l1_m_block = l1_task_offset // l1_num_n_blocks - l1_n_block = l1_task_offset % l1_num_n_blocks - l1_valid_m_blocks = T.ceildiv( - T.min(recv_counts[l1_expert], capacity), block_m - ) - if l1_m_block < l1_valid_m_blocks: - valid_rows = T.min( - block_m, - recv_counts[l1_expert] - l1_m_block * block_m, - ) - if tx == producer_begin: - T.wait_ge( - dispatch_arrivals[l1_expert, l1_m_block], - valid_rows, - scope=T.WaitScope.GPU, - semantics=T.WaitSemantics.ACQUIRE, - ) - T.sync_threads(5, producer_threads) - for k_block in T.serial(l1_num_k_blocks): - stage = (producer_step + k_block) % pipeline_stages - phase = ((producer_step + k_block) // pipeline_stages) & 1 - T.mbarrier_wait_parity( - stage_barriers[pipeline_stages + stage], phase ^ 1 - ) - T.tma_copy( - recv_x[ - l1_expert, - l1_m_block * block_m : (l1_m_block + 1) * block_m, - k_block * block_k : (k_block + 1) * block_k, - ], - a_shared[stage, :, :], - barrier=stage_barriers[stage], - ) - T.tma_copy( - l1_weight[ - l1_expert, - l1_n_block * block_n : (l1_n_block + 1) * block_n, - k_block * block_k : (k_block + 1) * block_k, - ], - b_shared[stage, :, :], - barrier=stage_barriers[stage], - ) - T.mbarrier_arrive(stage_barriers[stage]) - producer_step += l1_num_k_blocks - l1_task += num_sms - - for l2_round in T.serial(l2_total_rounds): - if l2_task < l2_total_tasks: - l2_task_offset = T.alloc_var(T.int32, init=l2_task) - l2_expert = l2_task_offset // (num_m_blocks * l2_num_n_blocks) - l2_task_offset -= l2_expert * num_m_blocks * l2_num_n_blocks - l2_m_block = l2_task_offset // l2_num_n_blocks - l2_n_block = l2_task_offset % l2_num_n_blocks - l2_valid_m_blocks = T.ceildiv( - T.min(recv_counts[l2_expert], capacity), block_m - ) - if l2_m_block < l2_valid_m_blocks: - if tx == producer_begin: - T.wait_ge( - l2_arrivals[l2_expert, l2_m_block], - l1_num_n_blocks, - scope=T.WaitScope.GPU, - semantics=T.WaitSemantics.ACQUIRE, - ) - T.sync_threads(6, producer_threads) - for k_block in T.serial(l2_num_k_blocks): - stage = (producer_step + k_block) % pipeline_stages - phase = ((producer_step + k_block) // pipeline_stages) & 1 - T.mbarrier_wait_parity( - stage_barriers[pipeline_stages + stage], phase ^ 1 - ) - T.tma_copy( - l2_x[ - l2_expert, - l2_m_block * block_m : (l2_m_block + 1) * block_m, - k_block * block_k : (k_block + 1) * block_k, - ], - a_shared[stage, :, :], - barrier=stage_barriers[stage], - ) - T.tma_copy( - l2_weight[ - l2_expert, - l2_n_block * block_n : (l2_n_block + 1) * block_n, - k_block * block_k : (k_block + 1) * block_k, - ], - b_shared[stage, :, :], - barrier=stage_barriers[stage], - ) - l2_a_sf_shared[stage, tx - producer_begin] = l2_x_sf[ - l2_expert, - l2_m_block * block_m + tx - producer_begin, - k_block, - ] - T.mbarrier_arrive(stage_barriers[stage]) - producer_step += l2_num_k_blocks - l2_task += num_sms - - elif tx >= math_begin and tx < math_end: - partial = T.alloc_fragment((block_m, block_n), T.float32) - accum = T.alloc_fragment((block_m, block_n), T.bfloat16) - gate = T.alloc_fragment((block_m, block_n // 2), T.float32) - gate_grouped = T.reshape( - gate, - ( - block_m, - num_output_scale_groups, - SCALE_GRANULARITY, - ), - ) - up = T.alloc_fragment((block_m, block_n // 2), T.float32) - amax = T.alloc_fragment( - (block_m, num_output_scale_groups), T.float32 - ) - scale = T.alloc_fragment( - (block_m, num_output_scale_groups), T.float32 - ) - quant_fp8 = T.alloc_fragment( - (block_m, block_n // 2), T.float8_e4m3fn - ) - act_scale = T.alloc_fragment((block_m,), T.float32) - weight_scale = T.alloc_local( - (2 * num_output_scale_groups,), T.float32 - ) - consumer_step = T.alloc_var(T.int32, init=0) - l1_task = T.alloc_var(T.int32, init=bid) - l2_task = T.alloc_var(T.int32, init=bid) - scatter_dst_rank = T.alloc_var(T.int32, init=0) - scatter_dst_token = T.alloc_var(T.int32, init=0) - scatter_dst_topk = T.alloc_var(T.int32, init=0) - - # Keep one surrounding scope while task-level completion is - # refined to block-level flags. - for expert_wave in T.serial(1): - - for l1_round in T.serial(l1_total_rounds): - if l1_task < l1_total_tasks: - tile_offset = T.alloc_var(T.int32, init=l1_task) - expert = tile_offset // (num_m_blocks * l1_num_n_blocks) - tile_offset -= expert * num_m_blocks * l1_num_n_blocks - m_block = tile_offset // l1_num_n_blocks - n_block = tile_offset % l1_num_n_blocks - l1_valid_m_blocks = T.ceildiv( - T.min(recv_counts[expert], capacity), block_m - ) - if m_block < l1_valid_m_blocks and ( - expert >= 0 - and l2_num_k_blocks * block_k - == intermediate_hidden - ): - T.clear(partial) - T.clear(accum) - for k_block in T.serial(l1_num_k_blocks): - stage = (consumer_step + k_block) % pipeline_stages - phase = ( - (consumer_step + k_block) - // pipeline_stages - ) & 1 - T.mbarrier_wait_parity( - stage_barriers[stage], phase - ) - T.gemm( - a_shared[stage, :, :], - b_shared[stage, :, :], - partial, - transpose_B=True, - ) - for i in T.Parallel(block_m): - act_scale[i] = recv_x_sf[ - expert, - m_block * block_m + i, - k_block, - ] - for scale_group in T.serial( - num_output_scale_groups - ): - weight_scale[2 * scale_group] = ( - l1_weight_sf[ - expert, - n_block - * num_output_scale_groups - + scale_group, - k_block, - ] - ) - weight_scale[2 * scale_group + 1] = ( - l1_weight_sf[ - expert, - num_l1_scale_groups - + n_block - * num_output_scale_groups - + scale_group, - k_block, - ] - ) - for i, j in T.Parallel(block_m, block_n): - accum[i, j] = ( - T.cast(partial[i, j], T.bfloat16) - * T.cast( - act_scale[i] - * weight_scale[ - 2 - * ( - j - // ( - 2 - * SCALE_GRANULARITY - ) - ) - + (j % 16) // 8 - ], - T.bfloat16, - ) - + accum[i, j] - ) - T.clear(partial) - T.mbarrier_arrive( - stage_barriers[ - pipeline_stages + stage - ] - ) - consumer_step += l1_num_k_blocks - - for i, j in T.Parallel( - block_m, block_n // 2 - ): - gate[i, j] = accum[ - i, (j // 8) * 16 + j % 8 - ] - for i, j in T.Parallel( - block_m, block_n // 2 - ): - up[i, j] = accum[ - i, (j // 8) * 16 + j % 8 + 8 - ] - for i, j in T.Parallel( - block_m, block_n // 2 - ): - clamped_gate = T.min( - gate[i, j], activation_clamp - ) - gate[i, j] = ( - clamped_gate - * T.sigmoid(clamped_gate) - * T.max( - T.min(up[i, j], activation_clamp), - -activation_clamp, - ) - * recv_weights[ - expert, m_block * block_m + i - ] - ) - T.reduce_absmax(gate_grouped, amax, dim=2) - for i, scale_group in T.Parallel( - block_m, num_output_scale_groups - ): - scale[i, scale_group] = ( - T.max(amax[i, scale_group], 1e-4) - / FP8_MAX - ) - l2_x_sf[ - expert, - m_block * block_m + i, - n_block - * num_output_scale_groups - + scale_group, - ] = scale[i, scale_group] - for i, j in T.Parallel( - block_m, block_n // 2 - ): - gate[i, j] = T.clamp( - gate[i, j] - / scale[ - i, j // SCALE_GRANULARITY - ], - -FP8_MAX, - FP8_MAX, - ) - T.copy(gate, quant_fp8) - T.copy(quant_fp8, l1_out_shared) - T.copy( - l1_out_shared, - l2_x[ - expert, - m_block - * block_m : (m_block + 1) - * block_m, - n_block - * (block_n // 2) : (n_block + 1) - * (block_n // 2), - ], - ) - if tx == math_begin: - T.atom_add( - l2_arrivals[expert, m_block], - 1, - scope="gpu", - sem="release", - ) - l1_task += num_sms - - for l2_round in T.serial(l2_total_rounds): - if l2_task < l2_total_tasks: - tile_offset = T.alloc_var(T.int32, init=l2_task) - expert = tile_offset // (num_m_blocks * l2_num_n_blocks) - tile_offset -= expert * num_m_blocks * l2_num_n_blocks - m_block = tile_offset // l2_num_n_blocks - n_block = tile_offset % l2_num_n_blocks - l2_valid_m_blocks = T.ceildiv( - T.min(recv_counts[expert], capacity), block_m - ) - if m_block < l2_valid_m_blocks and expert >= 0: - T.clear(partial) - T.clear(accum) - for k_block in T.serial(l2_num_k_blocks): - stage = (consumer_step + k_block) % pipeline_stages - phase = ( - (consumer_step + k_block) - // pipeline_stages - ) & 1 - T.mbarrier_wait_parity( - stage_barriers[stage], phase - ) - T.gemm( - a_shared[stage, :, :], - b_shared[stage, :, :], - partial, - transpose_B=True, - ) - for i in T.Parallel(block_m): - act_scale[i] = l2_a_sf_shared[ - stage, i - ] - weight_scale[0] = l2_weight_sf[ - expert, n_block * 2, k_block - ] - weight_scale[1] = l2_weight_sf[ - expert, n_block * 2 + 1, k_block - ] - for i, j in T.Parallel(block_m, block_n): - accum[i, j] = ( - T.cast(partial[i, j], T.bfloat16) - * T.cast( - act_scale[i] - * weight_scale[j // 128], - T.bfloat16, - ) - + accum[i, j] - ) - T.clear(partial) - T.mbarrier_arrive( - stage_barriers[ - pipeline_stages + stage - ] - ) - consumer_step += l2_num_k_blocks - T.copy(accum, l2_out_shared) - - scatter_warp = ( - tx - math_begin - ) // warp_size - for row_in_warp in T.serial(rows_per_math_warp): - row = ( - scatter_warp * rows_per_math_warp - + row_in_warp - ) - pool_row = m_block * block_m + row - if pool_row < recv_counts[expert]: - if tx % warp_size == 0: - scatter_dst_rank = src_ranks[ - expert, pool_row - ] - scatter_dst_token = src_tokens[ - expert, pool_row - ] - scatter_dst_topk = src_topk[ - expert, pool_row - ] - dst_rank = T.shfl_sync( - scatter_dst_rank, 0 - ) - dst_token = T.shfl_sync( - scatter_dst_token, 0 - ) - dst_topk = T.shfl_sync( - scatter_dst_topk, 0 - ) - if ( - dst_rank >= 0 - and dst_rank < num_ranks - and dst_token >= 0 - and dst_token < num_tokens - and dst_topk >= 0 - and dst_topk < num_topk - ): - T.put_warp( - T.address_of( - l2_out_shared[row, 0] - ), - T.address_of( - combine[ - dst_token, - dst_topk, - n_block * block_n, - ] - ), - block_n, - dst_pe=dst_rank, - unroll_factor=1, - ) - T.fence_sys() - T.sync_threads(4, num_math_threads) - if tx == math_begin: - for ready_rank in T.serial(num_ranks): - T.st( - l2_task_ready[expert, m_block, n_block], - 1, - scope="sys", - sem="release", - dst_pe=ready_rank, - ) - T.sync_threads(4, num_math_threads) - l2_task += num_sms - - - T.fence_sys() - T.sync_grid() - - return main - - -def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): - model_name, model = two_kernel.resolve_model_config(args) - hidden = model["hidden"] - intermediate_hidden = model["intermediate_hidden"] - num_experts = model["num_experts"] - num_topk = model["num_topk"] - num_tokens = args.num_tokens - activation_clamp = args.activation_clamp - - assert num_tokens > 0 - assert hidden >= 512 and hidden % 256 == 0 - assert intermediate_hidden > 0 and intermediate_hidden % 128 == 0 - assert num_experts > 0 and num_experts % num_local_ranks == 0 - assert 0 < num_topk <= min(32, num_experts) - num_experts_per_rank = num_experts // num_local_ranks - average_recv = ( - num_tokens * num_local_ranks * num_topk + num_experts - 1 - ) // num_experts - capacity = ( - args.capacity - if args.capacity is not None - else (max(average_recv * 2, 64) + 63) // 64 * 64 - ) - assert capacity >= 64 and capacity % 64 == 0 - - rank, num_ranks, group = init_dist(local_rank, num_local_ranks) - assert rank == local_rank and num_ranks == num_local_ranks - num_sms = torch.cuda.get_device_properties(local_rank).multi_processor_count - allocator = get_allocator( - size=two_kernel._allocator_size_bytes( - num_tokens, - hidden, - intermediate_hidden, - num_experts_per_rank, - num_topk, - capacity, - ), - device=f"cuda:{local_rank}", - is_distributed=True, - local_rank=local_rank, - num_local_ranks=num_local_ranks, - group=group, - use_vmm=True, - ) - - shape_family, config = select_single_kernel_config( - hidden, - intermediate_hidden, - num_tokens, - num_topk, - num_experts_per_rank, - num_sms, - ) - if args.experts_per_wave is not None: - assert args.experts_per_wave > 0 - assert num_experts_per_rank % args.experts_per_wave == 0 - config["num_experts_per_wave"] = args.experts_per_wave - if args.pipeline_stages is not None: - assert 2 <= args.pipeline_stages <= 4 - config["pipeline_stages"] = args.pipeline_stages - num_expert_waves = ( - num_experts_per_rank // config["num_experts_per_wave"] - ) - - spec = fused_single_kernel( - num_tokens, - hidden, - intermediate_hidden, - num_experts, - num_topk, - num_ranks, - capacity, - num_sms, - activation_clamp=activation_clamp, - **config, - ) - kernel = tilelang.compile(spec, compile_once=True, compile_group=group) - kernel.initialize(allocator=allocator) - if local_rank == 0 and args.print_source: - print(kernel.get_kernel_source()) - - torch.manual_seed(args.seed + local_rank) - x_bf16 = torch.randn( - (num_tokens, hidden), dtype=torch.bfloat16, device="cuda" - ) - x_fp8_src, x_sf_src = two_kernel.per_token_cast_to_fp8(x_bf16) - scores = torch.randn( - (num_tokens, num_experts), dtype=torch.float32, device="cuda" - ) - topk_weights_src, topk_idx_src = torch.topk( - scores, num_topk, dim=-1, sorted=False - ) - topk_idx_src = topk_idx_src.to(torch.int32) - - l1_bf16 = ( - torch.randn( - (num_experts_per_rank, 2 * intermediate_hidden, hidden), - dtype=torch.bfloat16, - device="cuda", - ) - * 0.05 - ) - l2_bf16 = ( - torch.randn( - (num_experts_per_rank, hidden, intermediate_hidden), - dtype=torch.bfloat16, - device="cuda", - ) - * 0.05 - ) - l1_fp8_src, l1_sf_src = two_kernel.block_cast_to_fp8(l1_bf16) - l2_fp8_src, l2_sf_src = two_kernel.block_cast_to_fp8(l2_bf16) - del scores, l1_bf16, l2_bf16 - - tensor = two_kernel.allocator_tensor - barrier = tensor((num_ranks,), torch.int32, allocator=allocator) - x = tensor(x_fp8_src.shape, x_fp8_src.dtype, allocator=allocator).copy_( - x_fp8_src - ) - x_sf = tensor(x_sf_src.shape, x_sf_src.dtype, allocator=allocator).copy_( - x_sf_src - ) - topk_idx = tensor( - topk_idx_src.shape, topk_idx_src.dtype, allocator=allocator - ).copy_(topk_idx_src) - topk_weights = tensor( - topk_weights_src.shape, topk_weights_src.dtype, allocator=allocator - ).copy_(topk_weights_src) - l1_fp8_kernel = two_kernel.interleave_gate_up_weights(l1_fp8_src) - l1_fp8 = tensor( - l1_fp8_kernel.shape, l1_fp8_kernel.dtype, allocator=allocator - ).copy_(l1_fp8_kernel) - del l1_fp8_kernel - l1_sf = tensor(l1_sf_src.shape, l1_sf_src.dtype, allocator=allocator).copy_( - l1_sf_src - ) - l2_fp8 = tensor( - l2_fp8_src.shape, l2_fp8_src.dtype, allocator=allocator - ).copy_(l2_fp8_src) - l2_sf = tensor(l2_sf_src.shape, l2_sf_src.dtype, allocator=allocator).copy_( - l2_sf_src - ) - - route_counts = tensor( - (num_ranks, num_experts), torch.int32, allocator=allocator - ) - recv_counts = tensor( - (num_experts_per_rank,), torch.int32, allocator=allocator - ) - recv_x = tensor( - (num_experts_per_rank, capacity, hidden), - torch.float8_e4m3fn, - allocator=allocator, - ) - recv_x_sf = tensor( - (num_experts_per_rank, capacity, hidden // SCALE_GRANULARITY), - torch.float32, - allocator=allocator, - ) - recv_weights = tensor( - (num_experts_per_rank, capacity), torch.float32, allocator=allocator - ) - src_ranks = tensor( - (num_experts_per_rank, capacity), torch.int32, allocator=allocator - ) - src_tokens = tensor( - (num_experts_per_rank, capacity), torch.int32, allocator=allocator - ) - src_topk = tensor( - (num_experts_per_rank, capacity), torch.int32, allocator=allocator - ) - route_slots = tensor( - (num_tokens, num_topk), torch.int32, allocator=allocator - ) - dispatch_arrivals = tensor( - (num_experts_per_rank, capacity // 64), - torch.uint32, - allocator=allocator, - ) - l2_arrivals = tensor( - (num_experts_per_rank, capacity // 64), - torch.uint32, - allocator=allocator, - ) - l2_task_ready = tensor( - ( - num_experts_per_rank, - capacity // config["block_m"], - hidden // config["block_n"], - ), - torch.uint32, - allocator=allocator, - ) - l2_x = tensor( - (num_experts_per_rank, capacity, intermediate_hidden), - torch.float8_e4m3fn, - allocator=allocator, - ) - l2_x_sf = tensor( - ( - num_experts_per_rank, - capacity, - intermediate_hidden // SCALE_GRANULARITY, - ), - torch.float32, - allocator=allocator, - ) - combine = tensor( - (num_tokens, num_topk, hidden), torch.bfloat16, allocator=allocator - ) - out = tensor((num_tokens, hidden), torch.bfloat16, allocator=allocator) - - def reset_state(): - route_counts.zero_() - barrier.zero_() - recv_counts.zero_() - dispatch_arrivals.zero_() - l2_arrivals.zero_() - l2_task_ready.zero_() - recv_x.zero_() - recv_x_sf.zero_() - recv_weights.zero_() - src_ranks.fill_(-1) - combine.zero_() - torch.cuda.synchronize() - dist.barrier(group=group) - - def run_pipeline(check_capacity: bool = False): - kernel( - x, - x_sf, - topk_idx, - topk_weights, - route_counts, - recv_counts, - route_slots, - dispatch_arrivals, - l2_arrivals, - l2_task_ready, - recv_x, - recv_x_sf, - recv_weights, - src_ranks, - src_tokens, - src_topk, - l1_fp8, - l1_sf, - l2_fp8, - l2_sf, - l2_x, - l2_x_sf, - combine, - barrier, - out, - ) - if check_capacity: - local_max = recv_counts.max() - dist.all_reduce(local_max, op=dist.ReduceOp.MAX, group=group) - assert local_max.item() <= capacity - return out - - reset_state() - actual = run_pipeline(check_capacity=True) - torch.cuda.synchronize() - dist.barrier(group=group) - - if args.check: - expected = two_kernel.torch_reference( - x_fp8_src, - x_sf_src, - topk_idx_src, - topk_weights_src, - l1_fp8_src, - l1_sf_src, - l2_fp8_src, - l2_sf_src, - group, - activation_clamp, - ) - diff = two_kernel.calc_diff(actual, expected) - assert diff < args.diff_tol, ( - f"rank {local_rank}: diff={diff} exceeds {args.diff_tol}" - ) - print(f"rank {local_rank} check passed, diff={diff:.6f}") - - if args.rep > 0: - reset_state() - for _ in range(args.warmup): - run_pipeline() - reset_state() - latency = do_bench( - run_pipeline, - warmup=0, - rep=args.rep, - post_fn=reset_state, - group=group, - ) - if local_rank == 0: - print( - "tilescale sm90 fp8 mega moe single kernel: " - f"model={model_name} family={shape_family} M={num_tokens} " - f"H={hidden} IH={intermediate_hidden} E={num_experts} " - f"topk={num_topk} capacity={capacity} " - f"epw={config['num_experts_per_wave']} " - f"stages={config['pipeline_stages']} " - f"latency={latency * 1000:.1f} us" - ) - - allocator.close() - dist.destroy_process_group() - - -if __name__ == "__main__": - parser = argparse.ArgumentParser() - parser.add_argument("--num-processes", type=int, default=8) - parser.add_argument( - "--model-config", choices=tuple(MODEL_CONFIGS), default="smoke" - ) - parser.add_argument("--hidden", type=int, default=None) - parser.add_argument("--intermediate-hidden", type=int, default=None) - parser.add_argument("--num-experts", type=int, default=None) - parser.add_argument("--num-topk", type=int, default=None) - parser.add_argument("--num-tokens", type=int, default=64) - parser.add_argument("--capacity", type=int, default=None) - parser.add_argument("--experts-per-wave", type=int, default=None) - parser.add_argument("--pipeline-stages", type=int, default=None) - parser.add_argument("--activation-clamp", type=float, default=10.0) - parser.add_argument("--seed", type=int, default=0) - parser.add_argument("--diff-tol", type=float, default=0.01) - parser.add_argument("--warmup", type=int, default=1) - parser.add_argument("--rep", type=int, default=1) - parser.add_argument("--check", action="store_true") - parser.add_argument("--print-source", action="store_true") - args = parser.parse_args() - torch.multiprocessing.spawn( - main, - args=(args.num_processes, args), - nprocs=args.num_processes, - ) From 2e8a973c96ddb19f681adea4fde743e30cb74fee Mon Sep 17 00:00:00 2001 From: zyy3077 Date: Sat, 29 Aug 2026 00:54:45 +0800 Subject: [PATCH 30/30] style(distributed): format SM90 MegaMoE example --- .../mega_moe/example_sm90_fp8_mega_moe.py | 578 ++++++++++-------- 1 file changed, 338 insertions(+), 240 deletions(-) diff --git a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py index 98efbe7fb1..b845e5152c 100644 --- a/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py +++ b/examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py @@ -35,6 +35,7 @@ FP8_MAX = 448.0 SCALE_GRANULARITY = 128 + def fused_l1_swiglu_manual_warp_kernel( num_tokens: int, hidden: int, @@ -147,7 +148,9 @@ def main( assign_topk = assign_route % num_topk assign_expert = topk_idx[assign_token, assign_topk] if assign_expert >= 0 and assign_expert < num_experts: - route_slots[assign_token, assign_topk] = T.atomic_add(route_counts[src_rank[0], assign_expert], 1, memory_order="relaxed", return_prev=True) + route_slots[assign_token, assign_topk] = T.atomic_add( + route_counts[src_rank[0], assign_expert], 1, memory_order="relaxed", return_prev=True + ) else: route_slots[assign_token, assign_topk] = -1 T.sync_threads(7, route_threads) @@ -155,7 +158,13 @@ def main( publish_warp = route_tid // warp_size for publish_rank in T.serial(publish_warp, num_ranks, route_threads // warp_size): if publish_rank != src_rank[0]: - T.put_warp(T.address_of(route_counts[src_rank[0], 0]), T.address_of(route_counts[src_rank[0], 0]), num_experts, dst_pe=publish_rank, unroll_factor=8) + T.put_warp( + T.address_of(route_counts[src_rank[0], 0]), + T.address_of(route_counts[src_rank[0], 0]), + num_experts, + dst_pe=publish_rank, + unroll_factor=8, + ) T.barrier_blocks(barrier[0]) @@ -208,10 +217,20 @@ def main( if metadata_expert >= 0 and metadata_slot >= 0 and metadata_slot < capacity: metadata_rank = metadata_expert // num_experts_per_rank metadata_local_expert = metadata_expert % num_experts_per_rank - T.st(recv_weights[metadata_local_expert, metadata_slot], topk_weights[metadata_token, metadata_topk], dst_pe=metadata_rank) + T.st( + recv_weights[metadata_local_expert, metadata_slot], + topk_weights[metadata_token, metadata_topk], + dst_pe=metadata_rank, + ) T.st(src_tokens[metadata_local_expert, metadata_slot], metadata_token, dst_pe=metadata_rank) T.st(src_topk[metadata_local_expert, metadata_slot], metadata_topk, dst_pe=metadata_rank) - T.st(src_ranks[metadata_local_expert, metadata_slot], src_rank[0], scope="sys", sem="release", dst_pe=metadata_rank) + T.st( + src_ranks[metadata_local_expert, metadata_slot], + src_rank[0], + scope="sys", + sem="release", + dst_pe=metadata_rank, + ) for pull_idx in T.serial(bid * dispatch_warps + dispatch_warp, num_m_tasks[0] * block_m, num_sms * dispatch_warps): pull_m_task = m_tasks[pull_idx // block_m] @@ -223,8 +242,20 @@ def main( T.sync_warp() pull_rank = src_ranks[pull_expert, pull_slot] pull_token = src_tokens[pull_expert, pull_slot] - T.get_warp(T.address_of(x[pull_token, 0]), T.address_of(recv_x[pull_expert, pull_slot, 0]), hidden, src_pe=pull_rank, unroll_factor=8) - T.get_warp(T.address_of(x_sf[pull_token, 0]), T.address_of(recv_x_sf[pull_expert, pull_slot, 0]), num_scale_groups, src_pe=pull_rank, unroll_factor=8) + T.get_warp( + T.address_of(x[pull_token, 0]), + T.address_of(recv_x[pull_expert, pull_slot, 0]), + hidden, + src_pe=pull_rank, + unroll_factor=8, + ) + T.get_warp( + T.address_of(x_sf[pull_token, 0]), + T.address_of(recv_x_sf[pull_expert, pull_slot, 0]), + num_scale_groups, + src_pe=pull_rank, + unroll_factor=8, + ) T.sync_warp() if dispatch_lane == dispatch_leader_lane: T.atom_add(arrivals[pull_expert, pull_slot // block_m], 1, scope="gpu", sem="release") @@ -238,49 +269,51 @@ def main( producer_m = producer_m_task % num_m_blocks producer_expert = producer_m_task // num_m_blocks if producer_n * block_n < l1_n: - producer_arrivals = T.min(block_m, recv_counts[producer_expert] - producer_m * block_m) - if tx == producer_begin: - T.wait_ge(arrivals[producer_expert, producer_m], producer_arrivals, scope=T.WaitScope.GPU, semantics=T.WaitSemantics.ACQUIRE) - T.sync_threads(5, producer_threads) - for producer_k in T.serial(num_k_blocks): - producer_stage = (producer_step + producer_k) % pipeline_stages - producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 - T.mbarrier_wait_parity(stage_barriers[pipeline_stages + producer_stage], producer_phase ^ 1) - producer_sf = T.alloc_local((1,), T.float32) - producer_sf[0] = recv_x_sf[ - producer_expert, producer_m * block_m + tx - producer_begin, producer_k] - for producer_ks in T.unroll(num_k_sub): - producer_sf_k = producer_k * num_k_sub + producer_ks + producer_arrivals = T.min(block_m, recv_counts[producer_expert] - producer_m * block_m) + if tx == producer_begin: + T.wait_ge( + arrivals[producer_expert, producer_m], + producer_arrivals, + scope=T.WaitScope.GPU, + semantics=T.WaitSemantics.ACQUIRE, + ) + T.sync_threads(5, producer_threads) + for producer_k in T.serial(num_k_blocks): + producer_stage = (producer_step + producer_k) % pipeline_stages + producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 + T.mbarrier_wait_parity(stage_barriers[pipeline_stages + producer_stage], producer_phase ^ 1) + producer_sf = T.alloc_local((1,), T.float32) + producer_sf[0] = recv_x_sf[producer_expert, producer_m * block_m + tx - producer_begin, producer_k] + for producer_ks in T.unroll(num_k_sub): + producer_sf_k = producer_k * num_k_sub + producer_ks + T.tma_copy( + recv_x[ + producer_expert, + producer_m * block_m : (producer_m + 1) * block_m, + producer_sf_k * SCALE_GRANULARITY : (producer_sf_k + 1) * SCALE_GRANULARITY, + ], + a_shared[producer_stage, producer_ks, :, :], + barrier=stage_barriers[producer_stage], + ) + for producer_n_block in T.serial(num_tma_n_blocks): T.tma_copy( - recv_x[ + l1_weight[ producer_expert, - producer_m * block_m : (producer_m + 1) * block_m, + producer_n * block_n + producer_n_block * tma_block_n : producer_n * block_n + + (producer_n_block + 1) * tma_block_n, producer_sf_k * SCALE_GRANULARITY : (producer_sf_k + 1) * SCALE_GRANULARITY, ], - a_shared[producer_stage, producer_ks, :, :], + b_shared[ + producer_stage, + producer_ks, + producer_n_block * tma_block_n : (producer_n_block + 1) * tma_block_n, + :, + ], barrier=stage_barriers[producer_stage], ) - for producer_n_block in T.serial(num_tma_n_blocks): - T.tma_copy( - l1_weight[ - producer_expert, - producer_n * block_n - + producer_n_block * tma_block_n : producer_n * block_n - + (producer_n_block + 1) * tma_block_n, - producer_sf_k * SCALE_GRANULARITY : (producer_sf_k + 1) * SCALE_GRANULARITY, - ], - b_shared[ - producer_stage, - producer_ks, - producer_n_block - * tma_block_n : (producer_n_block + 1) * tma_block_n, - :, - ], - barrier=stage_barriers[producer_stage], - ) - act_sf_shared[producer_stage, tx - producer_begin] = producer_sf[0] - T.mbarrier_arrive(stage_barriers[producer_stage]) - producer_step += num_k_blocks + act_sf_shared[producer_stage, tx - producer_begin] = producer_sf[0] + T.mbarrier_arrive(stage_barriers[producer_stage]) + producer_step += num_k_blocks else: T.inc_max_nreg(math_registers) @@ -304,72 +337,81 @@ def main( consumer_m = consumer_m_task % num_m_blocks consumer_expert = consumer_m_task // num_m_blocks if consumer_n * block_n < l1_n: - T.clear(partial) - T.clear(accum) - for consumer_k in T.serial(num_k_blocks): - consumer_stage = (consumer_step + consumer_k) % pipeline_stages - consumer_phase = ((consumer_step + consumer_k) // pipeline_stages) & 1 - T.mbarrier_wait_parity(stage_barriers[consumer_stage], consumer_phase) - # One TMA stage spans num_k_sub scale groups; WGMMA and promotion - # still run per SCALE_GRANULARITY so the per-128 scales stay exact. - for consumer_ks in T.unroll(num_k_sub): - consumer_sf_k = consumer_k * num_k_sub + consumer_ks - for scale_group in T.serial(num_output_scale_groups): - weight_scale[2 * scale_group] = l1_weight_sf[consumer_expert, consumer_n * num_output_scale_groups + scale_group, consumer_sf_k] - weight_scale[2 * scale_group + 1] = l1_weight_sf[consumer_expert, num_l1_scale_groups + consumer_n * num_output_scale_groups + scale_group, consumer_sf_k] - for i in T.Parallel(block_m): - act_scale[i] = act_sf_shared[consumer_stage, i] - T.gemm( - a_shared[consumer_stage, consumer_ks, :, :], - b_shared[consumer_stage, consumer_ks, :, :], - partial, transpose_B=True, clear_accum=True) - for i, j in T.Parallel(block_m, block_n): - accum[i, j] = partial[i, j] * ( - act_scale[i] * weight_scale[2 * (j // (2 * SCALE_GRANULARITY)) + (j % 16) // 8] - ) + accum[i, j] - T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) - consumer_step += num_k_blocks - for i, j in T.Parallel(block_m, block_n // 2): - gate[i, j] = accum[i, (j // 8) * 16 + j % 8] - for i, j in T.Parallel(block_m, block_n // 2): - up[i, j] = accum[i, (j // 8) * 16 + j % 8 + 8] - for i, j in T.Parallel(block_m, block_n // 2): - gate[i, j] = ( - T.min(gate[i, j], activation_clamp) - * T.sigmoid(T.min(gate[i, j], activation_clamp)) - * T.max( - T.min(up[i, j], activation_clamp), - -activation_clamp, - ) - * recv_weights[ + T.clear(partial) + T.clear(accum) + for consumer_k in T.serial(num_k_blocks): + consumer_stage = (consumer_step + consumer_k) % pipeline_stages + consumer_phase = ((consumer_step + consumer_k) // pipeline_stages) & 1 + T.mbarrier_wait_parity(stage_barriers[consumer_stage], consumer_phase) + # One TMA stage spans num_k_sub scale groups; WGMMA and promotion + # still run per SCALE_GRANULARITY so the per-128 scales stay exact. + for consumer_ks in T.unroll(num_k_sub): + consumer_sf_k = consumer_k * num_k_sub + consumer_ks + for scale_group in T.serial(num_output_scale_groups): + weight_scale[2 * scale_group] = l1_weight_sf[ + consumer_expert, consumer_n * num_output_scale_groups + scale_group, consumer_sf_k + ] + weight_scale[2 * scale_group + 1] = l1_weight_sf[ consumer_expert, - consumer_m * block_m + i, + num_l1_scale_groups + consumer_n * num_output_scale_groups + scale_group, + consumer_sf_k, ] + for i in T.Parallel(block_m): + act_scale[i] = act_sf_shared[consumer_stage, i] + T.gemm( + a_shared[consumer_stage, consumer_ks, :, :], + b_shared[consumer_stage, consumer_ks, :, :], + partial, + transpose_B=True, + clear_accum=True, ) - T.reduce_absmax(gate_grouped, amax, dim=2) - for i, scale_group in T.Parallel(block_m, num_output_scale_groups): - scale[i, scale_group] = T.max(amax[i, scale_group], 1e-4) / FP8_MAX - l2_x_sf[ + for i, j in T.Parallel(block_m, block_n): + accum[i, j] = ( + partial[i, j] * (act_scale[i] * weight_scale[2 * (j // (2 * SCALE_GRANULARITY)) + (j % 16) // 8]) + + accum[i, j] + ) + T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) + consumer_step += num_k_blocks + for i, j in T.Parallel(block_m, block_n // 2): + gate[i, j] = accum[i, (j // 8) * 16 + j % 8] + for i, j in T.Parallel(block_m, block_n // 2): + up[i, j] = accum[i, (j // 8) * 16 + j % 8 + 8] + for i, j in T.Parallel(block_m, block_n // 2): + gate[i, j] = ( + T.min(gate[i, j], activation_clamp) + * T.sigmoid(T.min(gate[i, j], activation_clamp)) + * T.max( + T.min(up[i, j], activation_clamp), + -activation_clamp, + ) + * recv_weights[ consumer_expert, consumer_m * block_m + i, - consumer_n * num_output_scale_groups + scale_group, - ] = scale[i, scale_group] - for i, j in T.Parallel(block_m, block_n // 2): - gate[i, j] = T.clamp(gate[i, j] / scale[i, j // SCALE_GRANULARITY], -FP8_MAX, FP8_MAX) - T.copy(gate, out_shared) - T.copy( - out_shared, - l2_x[ - consumer_expert, - consumer_m * block_m, - consumer_n * (block_n // 2), - ], + ] ) + T.reduce_absmax(gate_grouped, amax, dim=2) + for i, scale_group in T.Parallel(block_m, num_output_scale_groups): + scale[i, scale_group] = T.max(amax[i, scale_group], 1e-4) / FP8_MAX + l2_x_sf[ + consumer_expert, + consumer_m * block_m + i, + consumer_n * num_output_scale_groups + scale_group, + ] = scale[i, scale_group] + for i, j in T.Parallel(block_m, block_n // 2): + gate[i, j] = T.clamp(gate[i, j] / scale[i, j // SCALE_GRANULARITY], -FP8_MAX, FP8_MAX) + T.copy(gate, out_shared) + T.copy( + out_shared, + l2_x[ + consumer_expert, + consumer_m * block_m, + consumer_n * (block_n // 2), + ], + ) return main - def fused_l2_scatter_reduce_manual_warp_kernel( num_tokens: int, hidden: int, @@ -462,31 +504,33 @@ def main( and producer_n * block_n < hidden and num_k_blocks * block_k == intermediate_hidden ): - for producer_k in T.serial(num_k_blocks): - producer_stage = (producer_step + producer_k) % pipeline_stages - producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 - T.mbarrier_wait_parity(stage_barriers[pipeline_stages + producer_stage], producer_phase ^ 1) - T.tma_copy( - a[ - producer_expert, - producer_m * block_m : (producer_m + 1) * block_m, - producer_k * block_k : (producer_k + 1) * block_k, - ], - a_shared[producer_stage, :, :], - barrier=stage_barriers[producer_stage], - ) - T.tma_copy( - b[ - producer_expert, - producer_n * block_n : (producer_n + 1) * block_n, - producer_k * block_k : (producer_k + 1) * block_k, - ], - b_shared[producer_stage, :, :], - barrier=stage_barriers[producer_stage], - ) - a_sf_shared[producer_stage, tx - producer_begin] = a_sf[producer_expert, producer_m * block_m + tx - producer_begin, producer_k] - T.mbarrier_arrive(stage_barriers[producer_stage]) - producer_step += num_k_blocks + for producer_k in T.serial(num_k_blocks): + producer_stage = (producer_step + producer_k) % pipeline_stages + producer_phase = ((producer_step + producer_k) // pipeline_stages) & 1 + T.mbarrier_wait_parity(stage_barriers[pipeline_stages + producer_stage], producer_phase ^ 1) + T.tma_copy( + a[ + producer_expert, + producer_m * block_m : (producer_m + 1) * block_m, + producer_k * block_k : (producer_k + 1) * block_k, + ], + a_shared[producer_stage, :, :], + barrier=stage_barriers[producer_stage], + ) + T.tma_copy( + b[ + producer_expert, + producer_n * block_n : (producer_n + 1) * block_n, + producer_k * block_k : (producer_k + 1) * block_k, + ], + b_shared[producer_stage, :, :], + barrier=stage_barriers[producer_stage], + ) + a_sf_shared[producer_stage, tx - producer_begin] = a_sf[ + producer_expert, producer_m * block_m + tx - producer_begin, producer_k + ] + T.mbarrier_arrive(stage_barriers[producer_stage]) + producer_step += num_k_blocks elif tx >= math_begin: # WG1-2 run WGMMA and scatter their BF16 column pairs remotely. @@ -509,81 +553,96 @@ def main( and consumer_n * block_n < hidden and num_k_blocks * block_k == intermediate_hidden ): - T.clear(partial) - T.clear(accum) - for consumer_k in T.serial(num_k_blocks): - consumer_stage = (consumer_step + consumer_k) % pipeline_stages - consumer_phase = ((consumer_step + consumer_k) // pipeline_stages) & 1 - T.mbarrier_wait_parity(stage_barriers[consumer_stage], consumer_phase) - T.gemm(a_shared[consumer_stage, :, :], b_shared[consumer_stage, :, :], partial, transpose_B=True, clear_accum=True) - for i in T.Parallel(block_m): - act_scale[i] = a_sf_shared[consumer_stage, i] - weight_scale[0] = b_sf[consumer_expert, consumer_n * 2, consumer_k] - weight_scale[1] = b_sf[consumer_expert, consumer_n * 2 + 1, consumer_k] - for i, j in T.Parallel(block_m, block_n): - accum[i, j] = partial[i, j] * act_scale[i] * weight_scale[j // 128] + accum[i, j] - T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) - consumer_step += num_k_blocks - if use_put_warp_scatter: - # Stage the complete tile once, then let each math warp - # scatter eight rows with aligned 16-byte remote stores. - T.copy(accum, scatter_shared) - T.sync_threads(4, num_math_threads) - scatter_warp = (tx - math_begin) // warp_size - for row_in_warp in T.serial(block_m // (num_math_threads // warp_size)): - row = scatter_warp * (block_m // (num_math_threads // warp_size)) + row_in_warp - pool_row = consumer_m * block_m + row - if pool_row < recv_counts[consumer_expert]: - if tx % warp_size == 0: - scatter_dst_rank = src_ranks[consumer_expert, pool_row] - scatter_dst_token = src_tokens[consumer_expert, pool_row] - scatter_dst_topk = src_topk[consumer_expert, pool_row] - dst_rank = T.shfl_sync(scatter_dst_rank, 0) - dst_token = T.shfl_sync(scatter_dst_token, 0) - dst_topk = T.shfl_sync(scatter_dst_topk, 0) - if ( - dst_rank >= 0 - and dst_rank < num_ranks - and dst_token >= 0 - and dst_token < num_tokens - and dst_topk >= 0 - and dst_topk < num_topk - ): - T.put_warp( - T.address_of(scatter_shared[row, 0]), - T.address_of(combine[dst_token, dst_topk, consumer_n * block_n]), - block_n, - dst_pe=dst_rank, - unroll_factor=1, - ) - T.sync_threads(4, num_math_threads) - else: - # Direct scatter maps fragment owners to packed BF16 stores. - scatter_math_thread = tx - math_begin - scatter_wg = scatter_math_thread // warpgroup_size - scatter_warp_in_wg = (scatter_math_thread % warpgroup_size) // warp_size - scatter_lane = scatter_math_thread % warp_size - for scatter_row_half in T.serial(2): - row = scatter_warp_in_wg * 16 + scatter_row_half * 8 + scatter_lane // 4 - pool_row = consumer_m * block_m + row - if pool_row < recv_counts[consumer_expert]: + T.clear(partial) + T.clear(accum) + for consumer_k in T.serial(num_k_blocks): + consumer_stage = (consumer_step + consumer_k) % pipeline_stages + consumer_phase = ((consumer_step + consumer_k) // pipeline_stages) & 1 + T.mbarrier_wait_parity(stage_barriers[consumer_stage], consumer_phase) + T.gemm( + a_shared[consumer_stage, :, :], b_shared[consumer_stage, :, :], partial, transpose_B=True, clear_accum=True + ) + for i in T.Parallel(block_m): + act_scale[i] = a_sf_shared[consumer_stage, i] + weight_scale[0] = b_sf[consumer_expert, consumer_n * 2, consumer_k] + weight_scale[1] = b_sf[consumer_expert, consumer_n * 2 + 1, consumer_k] + for i, j in T.Parallel(block_m, block_n): + accum[i, j] = partial[i, j] * act_scale[i] * weight_scale[j // 128] + accum[i, j] + T.mbarrier_arrive(stage_barriers[pipeline_stages + consumer_stage]) + consumer_step += num_k_blocks + if use_put_warp_scatter: + # Stage the complete tile once, then let each math warp + # scatter eight rows with aligned 16-byte remote stores. + T.copy(accum, scatter_shared) + T.sync_threads(4, num_math_threads) + scatter_warp = (tx - math_begin) // warp_size + for row_in_warp in T.serial(block_m // (num_math_threads // warp_size)): + row = scatter_warp * (block_m // (num_math_threads // warp_size)) + row_in_warp + pool_row = consumer_m * block_m + row + if pool_row < recv_counts[consumer_expert]: + if tx % warp_size == 0: scatter_dst_rank = src_ranks[consumer_expert, pool_row] scatter_dst_token = src_tokens[consumer_expert, pool_row] scatter_dst_topk = src_topk[consumer_expert, pool_row] - if ( - scatter_dst_rank >= 0 - and scatter_dst_rank < num_ranks - and scatter_dst_token >= 0 - and scatter_dst_token < num_tokens - and scatter_dst_topk >= 0 - and scatter_dst_topk < num_topk - ): - for scatter_col_chunk in T.serial(block_n // math_warpgroups // 8): - scatter_col = scatter_wg * (block_n // math_warpgroups) + scatter_col_chunk * 8 + (scatter_lane % 4) * 2 - scatter_value_lo = T.alloc_var(T.uint16, init=T.reinterpret(T.cast(accum[row, scatter_col], T.bfloat16), T.uint16)) - scatter_value_hi = T.alloc_var(T.uint16, init=T.reinterpret(T.cast(accum[row, scatter_col + 1], T.bfloat16), T.uint16)) - scatter_value = T.alloc_var(T.uint32, init=T.cast(scatter_value_lo, T.uint32) | (T.cast(scatter_value_hi, T.uint32) << 16)) - T.st(combine[scatter_dst_token, scatter_dst_topk, consumer_n * block_n + scatter_col], scatter_value, dst_pe=scatter_dst_rank) + dst_rank = T.shfl_sync(scatter_dst_rank, 0) + dst_token = T.shfl_sync(scatter_dst_token, 0) + dst_topk = T.shfl_sync(scatter_dst_topk, 0) + if ( + dst_rank >= 0 + and dst_rank < num_ranks + and dst_token >= 0 + and dst_token < num_tokens + and dst_topk >= 0 + and dst_topk < num_topk + ): + T.put_warp( + T.address_of(scatter_shared[row, 0]), + T.address_of(combine[dst_token, dst_topk, consumer_n * block_n]), + block_n, + dst_pe=dst_rank, + unroll_factor=1, + ) + T.sync_threads(4, num_math_threads) + else: + # Direct scatter maps fragment owners to packed BF16 stores. + scatter_math_thread = tx - math_begin + scatter_wg = scatter_math_thread // warpgroup_size + scatter_warp_in_wg = (scatter_math_thread % warpgroup_size) // warp_size + scatter_lane = scatter_math_thread % warp_size + for scatter_row_half in T.serial(2): + row = scatter_warp_in_wg * 16 + scatter_row_half * 8 + scatter_lane // 4 + pool_row = consumer_m * block_m + row + if pool_row < recv_counts[consumer_expert]: + scatter_dst_rank = src_ranks[consumer_expert, pool_row] + scatter_dst_token = src_tokens[consumer_expert, pool_row] + scatter_dst_topk = src_topk[consumer_expert, pool_row] + if ( + scatter_dst_rank >= 0 + and scatter_dst_rank < num_ranks + and scatter_dst_token >= 0 + and scatter_dst_token < num_tokens + and scatter_dst_topk >= 0 + and scatter_dst_topk < num_topk + ): + for scatter_col_chunk in T.serial(block_n // math_warpgroups // 8): + scatter_col = ( + scatter_wg * (block_n // math_warpgroups) + scatter_col_chunk * 8 + (scatter_lane % 4) * 2 + ) + scatter_value_lo = T.alloc_var( + T.uint16, init=T.reinterpret(T.cast(accum[row, scatter_col], T.bfloat16), T.uint16) + ) + scatter_value_hi = T.alloc_var( + T.uint16, init=T.reinterpret(T.cast(accum[row, scatter_col + 1], T.bfloat16), T.uint16) + ) + scatter_value = T.alloc_var( + T.uint32, + init=T.cast(scatter_value_lo, T.uint32) | (T.cast(scatter_value_hi, T.uint32) << 16), + ) + T.st( + combine[scatter_dst_token, scatter_dst_topk, consumer_n * block_n + scatter_col], + scatter_value, + dst_pe=scatter_dst_rank, + ) T.fence_sys() T.sync_grid() @@ -610,8 +669,8 @@ def main( reduce_n * reduce_block_h, ], ) - return main + return main def _allocator_size_bytes( @@ -627,29 +686,19 @@ def _allocator_size_bytes( fp32 = 4 i32 = 4 weight_bytes = num_experts_per_rank * (2 * intermediate_hidden * hidden * fp8 + hidden * intermediate_hidden * fp8) - weight_scale_bytes = num_experts_per_rank * ((2 * intermediate_hidden // 128) * (hidden // 128) + (hidden // 128) * (intermediate_hidden // 128)) * fp32 + weight_scale_bytes = ( + num_experts_per_rank * ((2 * intermediate_hidden // 128) * (hidden // 128) + (hidden // 128) * (intermediate_hidden // 128)) * fp32 + ) pool_bytes = ( num_experts_per_rank * capacity - * ( - hidden * fp8 - + (hidden // 128) * fp32 - + 4 * i32 - + intermediate_hidden * fp8 - + (intermediate_hidden // 128) * fp32 - ) - ) - input_bytes = num_tokens * ( - hidden * fp8 - + (hidden // 128) * fp32 - + num_topk * (3 * i32 + fp32) - + (num_topk + 1) * hidden * bf16 + * (hidden * fp8 + (hidden // 128) * fp32 + 4 * i32 + intermediate_hidden * fp8 + (intermediate_hidden // 128) * fp32) ) + input_bytes = num_tokens * (hidden * fp8 + (hidden // 128) * fp32 + num_topk * (3 * i32 + fp32) + (num_topk + 1) * hidden * bf16) total_bytes = weight_bytes + weight_scale_bytes + pool_bytes + input_bytes + 2**27 return (total_bytes + 2**20 - 1) // 2**20 * 2**20 - def torch_reference( x_fp8: torch.Tensor, x_sf: torch.Tensor, @@ -666,14 +715,13 @@ def _gather_cat(tensor: torch.Tensor, group: dist.ProcessGroup) -> torch.Tensor: gathered = [torch.empty_like(tensor) for _ in range(dist.get_world_size(group))] dist.all_gather(gathered, tensor, group=group) return torch.cat(gathered, dim=0) + l1_all = _gather_cat(l1_fp8, group) l1_sf_all = _gather_cat(l1_sf, group) l2_all = _gather_cat(l2_fp8, group) l2_sf_all = _gather_cat(l2_sf, group) x_m, x_k = x_fp8.shape - x = ( - x_fp8.float().view(x_m, x_k // SCALE_GRANULARITY, SCALE_GRANULARITY) * x_sf.unsqueeze(-1) - ).view(x_m, x_k) + x = (x_fp8.float().view(x_m, x_k // SCALE_GRANULARITY, SCALE_GRANULARITY) * x_sf.unsqueeze(-1)).view(x_m, x_k) result = torch.zeros((x.size(0), l2_all.size(1)), dtype=torch.float32, device=x.device) for expert_idx in range(l1_all.size(0)): @@ -688,9 +736,7 @@ def _gather_cat(tensor: torch.Tensor, group: dist.ProcessGroup) -> torch.Tensor: l1_weight_view = l1_weight_fp8.float().view( groups, l1_n // SCALE_GRANULARITY, SCALE_GRANULARITY, l1_k // SCALE_GRANULARITY, SCALE_GRANULARITY ) - l1_weight = ( - l1_weight_view * l1_weight_sf.unsqueeze(-1).unsqueeze(-3) - ).view(groups, l1_n, l1_k)[0] + l1_weight = (l1_weight_view * l1_weight_sf.unsqueeze(-1).unsqueeze(-3)).view(groups, l1_n, l1_k)[0] gate_up = x[token_indices] @ l1_weight.T gate, up = gate_up.chunk(2, dim=-1) gate = gate.clamp(max=activation_clamp) @@ -698,30 +744,23 @@ def _gather_cat(tensor: torch.Tensor, group: dist.ProcessGroup) -> torch.Tensor: activated = torch.nn.functional.silu(gate) * up activated *= topk_weights[token_indices, topk_slots].unsqueeze(-1) activated_m, activated_k = activated.shape - activated_view = activated.float().view( - activated_m, activated_k // SCALE_GRANULARITY, SCALE_GRANULARITY - ) + activated_view = activated.float().view(activated_m, activated_k // SCALE_GRANULARITY, SCALE_GRANULARITY) activated_sf = activated_view.abs().amax(dim=-1).clamp(1e-4) / FP8_MAX activated_fp8 = (activated_view / activated_sf.unsqueeze(-1)).to(torch.float8_e4m3fn) - activated_dequant = ( - activated_fp8.float() * activated_sf.unsqueeze(-1) - ).view(activated_m, activated_k) + activated_dequant = (activated_fp8.float() * activated_sf.unsqueeze(-1)).view(activated_m, activated_k) l2_weight_fp8 = l2_all[expert_idx : expert_idx + 1] l2_weight_sf = l2_sf_all[expert_idx : expert_idx + 1] groups, l2_n, l2_k = l2_weight_fp8.shape l2_weight_view = l2_weight_fp8.float().view( groups, l2_n // SCALE_GRANULARITY, SCALE_GRANULARITY, l2_k // SCALE_GRANULARITY, SCALE_GRANULARITY ) - l2_weight = ( - l2_weight_view * l2_weight_sf.unsqueeze(-1).unsqueeze(-3) - ).view(groups, l2_n, l2_k)[0] + l2_weight = (l2_weight_view * l2_weight_sf.unsqueeze(-1).unsqueeze(-3)).view(groups, l2_n, l2_k)[0] contribution = (activated_dequant @ l2_weight.T).to(torch.bfloat16).float() result.index_add_(0, token_indices, contribution) return result.to(torch.bfloat16) - def main(local_rank: int, num_local_ranks: int, args: argparse.Namespace): def per_token_cast_to_fp8(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: m, k = x.shape @@ -770,12 +809,14 @@ def select_manual_warp_configs( l1_stages = 3 l2_stages = 3 generic_experts_per_wave = num_experts_per_rank + def normalize_experts_per_wave(num_experts: int, requested: int) -> int: requested = min(max(requested, 1), num_experts) for candidate in range(requested, num_experts + 1): if num_experts % candidate == 0: return candidate return num_experts + if num_experts_per_rank <= routed_tokens <= 4 * num_experts_per_rank: expected_tokens = (routed_tokens + num_experts_per_rank - 1) // num_experts_per_rank num_m_blocks = (expected_tokens + 63) // 64 @@ -849,16 +890,34 @@ def normalize_experts_per_wave(num_experts: int, requested: int) -> int: use_vmm=True, ) - shape_family, l1_config, l2_config = select_manual_warp_configs(hidden, intermediate_hidden, num_tokens, num_topk, num_experts_per_rank, num_sms) + shape_family, l1_config, l2_config = select_manual_warp_configs( + hidden, intermediate_hidden, num_tokens, num_topk, num_experts_per_rank, num_sms + ) # Packed warp scatter amortizes its shared-memory staging for larger token batches. use_put_warp_scatter = num_tokens >= 1024 kernel_specs = [ fused_l1_swiglu_manual_warp_kernel( - num_tokens, hidden, 2 * intermediate_hidden, num_experts, num_topk, num_ranks, capacity, num_sms, - activation_clamp=activation_clamp, **l1_config, + num_tokens, + hidden, + 2 * intermediate_hidden, + num_experts, + num_topk, + num_ranks, + capacity, + num_sms, + activation_clamp=activation_clamp, + **l1_config, ), fused_l2_scatter_reduce_manual_warp_kernel( - num_tokens, hidden, intermediate_hidden, num_experts_per_rank, num_topk, num_ranks, capacity, num_sms, **l2_config, + num_tokens, + hidden, + intermediate_hidden, + num_experts_per_rank, + num_topk, + num_ranks, + capacity, + num_sms, + **l2_config, use_put_warp_scatter=use_put_warp_scatter, ), ] @@ -920,7 +979,9 @@ def allocator_tensor(shape, dtype, allocator): m_tasks = allocator_tensor((num_experts_per_rank * ((capacity + 63) // 64),), torch.int32, allocator=allocator) num_m_tasks = allocator_tensor((1,), torch.int32, allocator=allocator) l2_x = allocator_tensor((num_experts_per_rank, capacity, intermediate_hidden), torch.float8_e4m3fn, allocator=allocator) - l2_x_sf = allocator_tensor((num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), torch.float32, allocator=allocator) + l2_x_sf = allocator_tensor( + (num_experts_per_rank, capacity, intermediate_hidden // SCALE_GRANULARITY), torch.float32, allocator=allocator + ) combine = allocator_tensor((num_tokens, num_topk, hidden), torch.bfloat16, allocator=allocator) out = allocator_tensor((num_tokens, hidden), torch.bfloat16, allocator=allocator) @@ -939,17 +1000,46 @@ def reset_state(): def run_pipeline(check_capacity: bool = False): fused_l1( - x, x_sf, topk_idx, topk_weights, route_counts, recv_counts, route_slots, arrivals, - m_tasks, num_m_tasks, recv_x, recv_x_sf, recv_weights, src_ranks, src_tokens, src_topk, - l1_fp8, l1_sf, l2_x, l2_x_sf, barrier, + x, + x_sf, + topk_idx, + topk_weights, + route_counts, + recv_counts, + route_slots, + arrivals, + m_tasks, + num_m_tasks, + recv_x, + recv_x_sf, + recv_weights, + src_ranks, + src_tokens, + src_topk, + l1_fp8, + l1_sf, + l2_x, + l2_x_sf, + barrier, ) if check_capacity: local_max = recv_counts.max() dist.all_reduce(local_max, op=dist.ReduceOp.MAX, group=group) assert local_max.item() <= capacity, f"expert capacity {capacity} is smaller than received routes {local_max.item()}" fused_l2( - l2_x, l2_fp8, l2_x_sf, l2_sf, recv_counts, m_tasks, num_m_tasks, src_ranks, - src_tokens, src_topk, combine, barrier, out, + l2_x, + l2_fp8, + l2_x_sf, + l2_sf, + recv_counts, + m_tasks, + num_m_tasks, + src_ranks, + src_tokens, + src_topk, + combine, + barrier, + out, ) return out @@ -965,8 +1055,16 @@ def calc_diff(x: torch.Tensor, y: torch.Tensor) -> float: return (1 - 2 * (x * y).sum() / (x.square() + y.square()).sum()).item() expected = torch_reference( - x_fp8_src, x_sf_src, topk_idx_src, topk_weights_src, l1_fp8_src, - l1_sf_src, l2_fp8_src, l2_sf_src, group, activation_clamp, + x_fp8_src, + x_sf_src, + topk_idx_src, + topk_weights_src, + l1_fp8_src, + l1_sf_src, + l2_fp8_src, + l2_sf_src, + group, + activation_clamp, ) diff = calc_diff(actual, expected) assert diff < args.diff_tol, f"rank {local_rank}: diff={diff} exceeds {args.diff_tol}"