From 77940dec80c935bb38e9f55f03a054fc7cbc7aec Mon Sep 17 00:00:00 2001 From: Yuan Luo Date: Mon, 24 Aug 2026 13:37:14 +0800 Subject: [PATCH] [MoE] Gather the cutlass MoE activation and its scales in one launch (#34915) Co-authored-by: luoyuan.luo --- python/sglang/kernels/ops/moe/__init__.py | 13 ++ .../ops/moe/shuffle_rows_with_scales.py | 137 ++++++++++++++++++ python/sglang/srt/layers/moe/cutlass_moe.py | 10 +- .../moe/bench_shuffle_rows_with_scales.py | 69 +++++++++ .../ops/moe/test_shuffle_rows_with_scales.py | 123 ++++++++++++++++ 5 files changed, 350 insertions(+), 2 deletions(-) create mode 100644 python/sglang/kernels/ops/moe/shuffle_rows_with_scales.py create mode 100644 test/registered/kernels/benchmark/moe/bench_shuffle_rows_with_scales.py create mode 100644 test/registered/kernels/ops/moe/test_shuffle_rows_with_scales.py diff --git a/python/sglang/kernels/ops/moe/__init__.py b/python/sglang/kernels/ops/moe/__init__.py index ad7a62ad2..51e40eea3 100644 --- a/python/sglang/kernels/ops/moe/__init__.py +++ b/python/sglang/kernels/ops/moe/__init__.py @@ -201,3 +201,16 @@ register_kernel( description="MoE align-block-size, single-launch triton variant.", ) ) + +# One gather for a quantized activation and its group scales: replaces the pair +# of shuffle_rows launches the cutlass fp8 blockwise MoE used to walk the same +# dst2src map with. Byte-identical to those calls. +register_kernel( + KernelSpec( + op="moe.shuffle_rows_with_scales", + backend=KernelBackend.TRITON, + target="sglang.kernels.ops.moe.shuffle_rows_with_scales:shuffle_rows_with_scales", + capabilities=_CUDA, + description="Row gather of quantized values plus their scales, one launch.", + ) +) diff --git a/python/sglang/kernels/ops/moe/shuffle_rows_with_scales.py b/python/sglang/kernels/ops/moe/shuffle_rows_with_scales.py new file mode 100644 index 000000000..5bc616699 --- /dev/null +++ b/python/sglang/kernels/ops/moe/shuffle_rows_with_scales.py @@ -0,0 +1,137 @@ +"""Single-launch row gather for a quantized activation and its group scales. + +The cutlass fp8 blockwise MoE quantizes its activation once and then replicates +rows per routed expert. That took two ``shuffle_rows`` launches walking the same +dst2src map: one for the fp8 values, one for the fp32 group scales. The scale +gather moves 1/32 of the bytes the value gather does (``k // 128`` fp32 against +``k`` fp8), so as its own launch it is almost pure latency -- which is exactly +the cost that matters at low concurrency, where the whole gather is a few tens +of KB. This kernel walks the map once and writes both. + +The gather is a permutation of bytes -- rows are copied, never recomputed -- so +the result is bit-identical to the two calls it replaces. +""" + +from typing import Tuple + +import torch +import triton +import triton.language as tl + +# Bytes of the value row one program copies; the grid is +# (num_dst_rows, ceil(k / BLOCK_K)). +# +# This is a bytes-per-thread knob, not a parallelism knob, and that is what +# makes it load-bearing. At low concurrency the gather is a few tens of KB and +# every setting measures the same, because all that is being timed is the launch. +# At prefill sizes it decides everything: on B200 with k = 7168, rows = 8192, +# against the two shuffle_rows launches this replaces (33.9 us) -- +# +# BLOCK_K 512 1024 2048 4096 8192 16384 +# time 66.6 37.8 25.5 18.1 17.4 19.4 us (num_warps=4) +# +# 512 is half the speed of the CUDA kernel it replaces: at num_warps=4 that is +# 4 bytes per thread, a quarter of the 128 bits per thread the CUDA kernel +# vectorizes to. 4096 puts 32 bytes in each thread and lands on the plateau. +# +# Columns past k are masked off, so a model narrower than BLOCK_K runs partly +# empty lanes: k = 2048 still measures 1.33x against the two launches, and +# nothing narrower has been measured. If a k of 1024 or less turns up on this +# path, re-run the sweep before assuming this setting still holds. +BLOCK_K = 4096 +NUM_WARPS = 4 + + +@triton.jit +def _shuffle_rows_with_scales_kernel( + q_ptr, # [num_src_rows, k] int8 view of the quantized values + scale_ptr, # [num_src_rows, num_groups] fp32 group scales + q_out_ptr, # [num_dst_rows, k] int8 view + scale_out_ptr, # [num_dst_rows, num_groups] fp32 + dst2src_ptr, # [num_dst_rows] int32, out[i] = src[dst2src[i]] + k, + num_groups, + BLOCK_K: tl.constexpr, + BLOCK_G: tl.constexpr, +): + dst_row = tl.program_id(0) + tile = tl.program_id(1) + # int64 row bases: rows * k overflows int32 well inside the shapes this path + # serves (the CUDA shuffle_rows it replaces indexes in int64 for the same + # reason). + src_row = tl.load(dst2src_ptr + dst_row).to(tl.int64) + dst_row64 = dst_row.to(tl.int64) + + offs_k = tile * BLOCK_K + tl.arange(0, BLOCK_K) + mask_k = offs_k < k + vals = tl.load(q_ptr + src_row * k + offs_k, mask=mask_k) + tl.store(q_out_ptr + dst_row64 * k + offs_k, vals, mask=mask_k) + + # The scale row is 1/32 of the value row, so one tile carries all of it + # rather than the whole thing costing a second launch. + if tile == 0: + offs_g = tl.arange(0, BLOCK_G) + mask_g = offs_g < num_groups + scales = tl.load(scale_ptr + src_row * num_groups + offs_g, mask=mask_g) + tl.store(scale_out_ptr + dst_row64 * num_groups + offs_g, scales, mask=mask_g) + + +def shuffle_rows_with_scales( + q: torch.Tensor, + scale: torch.Tensor, + dst2src_map: torch.Tensor, + num_dst_rows: int, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Gather ``num_dst_rows`` rows of ``q`` and ``scale`` through one map. + + Replaces a pair of ``shuffle_rows`` calls over the same ``dst2src_map``, + with the same semantics for both tensors: ``out[i] = src[dst2src_map[i]]``. + Returns the two gathered tensors, allocated here. + + ``q`` is any 1-byte dtype (it is moved as bytes, not interpreted) and + ``scale`` is its row-major per-group scale tensor; both must be contiguous + and share a row count. + """ + assert q.dim() == 2 and scale.dim() == 2, "q and scale must be 2D" + assert q.is_contiguous() and scale.is_contiguous(), "q and scale must be contiguous" + assert q.element_size() == 1, f"q must be a 1-byte dtype, got {q.dtype}" + assert ( + q.shape[0] == scale.shape[0] + ), f"row count mismatch: q {q.shape[0]} vs scale {scale.shape[0]}" + assert ( + dst2src_map.numel() >= num_dst_rows + ), f"map holds {dst2src_map.numel()} rows, need {num_dst_rows}" + # The kernel reads the map as whatever dtype it carries and casts to int64, + # so a float map would truncate into a plausible-looking row id instead of + # failing. + assert dst2src_map.dtype in ( + torch.int32, + torch.int64, + ), f"dst2src_map must hold integer row ids, got {dst2src_map.dtype}" + assert q.device == scale.device == dst2src_map.device, ( + f"inputs must share a device: q {q.device}, scale {scale.device}, " + f"map {dst2src_map.device}" + ) + + k = q.shape[1] + num_groups = scale.shape[1] + q_out = torch.empty((num_dst_rows, k), device=q.device, dtype=q.dtype) + scale_out = torch.empty( + (num_dst_rows, num_groups), device=scale.device, dtype=scale.dtype + ) + if num_dst_rows == 0: + return q_out, scale_out + + _shuffle_rows_with_scales_kernel[(num_dst_rows, triton.cdiv(k, BLOCK_K))]( + q.view(torch.int8), + scale, + q_out.view(torch.int8), + scale_out, + dst2src_map, + k, + num_groups, + BLOCK_K=BLOCK_K, + BLOCK_G=triton.next_power_of_2(max(num_groups, 1)), + num_warps=NUM_WARPS, + ) + return q_out, scale_out diff --git a/python/sglang/srt/layers/moe/cutlass_moe.py b/python/sglang/srt/layers/moe/cutlass_moe.py index afc636fb1..59ddf6f40 100755 --- a/python/sglang/srt/layers/moe/cutlass_moe.py +++ b/python/sglang/srt/layers/moe/cutlass_moe.py @@ -19,6 +19,9 @@ if _is_cuda: ) from sglang.kernels.ops.activation.activation import silu_and_mul + from sglang.kernels.ops.moe.shuffle_rows_with_scales import ( + shuffle_rows_with_scales, + ) def cutlass_fused_experts_fp8( @@ -207,8 +210,11 @@ def cutlass_fused_experts_fp8( ) else: a_q, a1_scale = sglang_per_token_group_quant_fp8(a, 128) - rep_a_q = shuffle_rows(a_q, a_map, (m * topk, k)) - rep_a1_scales = shuffle_rows(a1_scale, a_map, (m * topk, int(k / 128))) + # One gather for both: the scale rows are 1/32 of the value rows, so + # walking the map a second time for them was almost pure launch latency. + rep_a_q, rep_a1_scales = shuffle_rows_with_scales( + a_q, a1_scale, a_map, m * topk + ) c1 = torch.empty((m * topk, n * 2), device=device, dtype=out_dtype) c2 = torch.empty((m * topk, k), device=device, dtype=out_dtype) diff --git a/test/registered/kernels/benchmark/moe/bench_shuffle_rows_with_scales.py b/test/registered/kernels/benchmark/moe/bench_shuffle_rows_with_scales.py new file mode 100644 index 000000000..daff5d63d --- /dev/null +++ b/test/registered/kernels/benchmark/moe/bench_shuffle_rows_with_scales.py @@ -0,0 +1,69 @@ +import torch +import triton +import triton.testing + +from sglang.kernels.jit.benchmark.utils import ( + DEFAULT_DEVICE, + get_benchmark_range, + run_benchmark, +) +from sglang.kernels.ops.moe.shuffle_rows_with_scales import shuffle_rows_with_scales +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci( + est_time=15, stage="base-b-kernel-benchmark", runner_config="1-gpu-large" +) + +GROUP_SIZE = 128 + +# (hidden, tokens). The cutlass fp8 blockwise MoE gathers tokens * topk rows out +# of a tokens-row source, so the small end is bs=1 decode and the large end is a +# prefill-sized batch. 7168 is a DeepSeek-class hidden size. +SHAPES = get_benchmark_range( + full_range=[(7168, 1), (2048, 1), (7168, 8), (7168, 64), (7168, 1024)], + ci_range=[(7168, 1), (7168, 1024)], +) + + +@triton.testing.perf_report( + triton.testing.Benchmark( + x_names=["k", "tokens"], + x_vals=SHAPES, + line_arg="provider", + line_vals=["fused", "two_launches"], + line_names=["Fused gather", "Two shuffle_rows"], + styles=[("blue", "-"), ("red", "--")], + ylabel="us", + plot_name="shuffle-rows-with-scales-performance", + args={"topk": 8}, + ) +) +def benchmark(k: int, tokens: int, topk: int, provider: str): + from sgl_kernel import shuffle_rows + + rows = tokens * topk + q = torch.randint( + 0, 256, (tokens, k), dtype=torch.uint8, device=DEFAULT_DEVICE + ).view(torch.float8_e4m3fn) + scale = torch.randn( + (tokens, k // GROUP_SIZE), dtype=torch.float32, device=DEFAULT_DEVICE + ) + # Duplicate source rows are the normal case: a token is replicated once per + # expert it routes to. + dst2src = torch.randint( + 0, tokens, (rows,), dtype=torch.int32, device=DEFAULT_DEVICE + ) + + if provider == "fused": + fn = lambda: shuffle_rows_with_scales(q, scale, dst2src, rows) + else: + fn = lambda: ( + shuffle_rows(q, dst2src, (rows, k)), + shuffle_rows(scale, dst2src, (rows, k // GROUP_SIZE)), + ) + + return run_benchmark(fn) + + +if __name__ == "__main__": + benchmark.run(print_data=True) diff --git a/test/registered/kernels/ops/moe/test_shuffle_rows_with_scales.py b/test/registered/kernels/ops/moe/test_shuffle_rows_with_scales.py new file mode 100644 index 000000000..0b6d46a6d --- /dev/null +++ b/test/registered/kernels/ops/moe/test_shuffle_rows_with_scales.py @@ -0,0 +1,123 @@ +"""Bit-exactness of the fused value+scale row gather. + +The oracle is plain torch advanced indexing, which is the whole contract: +``out[i] = src[dst2src_map[i]]`` for both tensors. A second test cross-checks the +pair of `shuffle_rows` calls this replaces, to back the drop-in claim. + +Comparisons are made on integer views. The values are fp8 and a random byte +pattern is a NaN often enough that `torch.equal` on the float view would report a +difference where the bytes agree -- and bytes are exactly what this kernel +promises to preserve. +""" + +import itertools +import sys + +import pytest +import torch + +from sglang.kernels.jit.utils import get_ci_test_range +from sglang.kernels.ops.moe.shuffle_rows_with_scales import shuffle_rows_with_scales +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large") + +GROUP_SIZE = 128 + +# The CUDA shuffle_rows this replaces loads 128 bits per thread and takes its +# element count as num_cols / elems_per_thread with no remainder handling, so it +# only moves a whole fp32 scale row when (k // GROUP_SIZE) % 4 == 0. Shapes below +# that bound are checked against the torch oracle only -- the reference itself +# drops the tail there. +CUDA_XCHECK_SCALE_ALIGN = 4 + +# k = 896 gives 7 scale groups, the shape whose scale tail the CUDA reference +# cannot express; 7168 is the DeepSeek-class hidden size. Destination rows exceed +# source rows because the map replicates each token once per routed expert. +CASES = get_ci_test_range( + [ + (k, src, dst) + for k, (src, dst) in itertools.product( + [512, 896, 2560, 7168], [(1, 8), (17, 136), (64, 512)] + ) + ], + [ + (896, 1, 8), + (7168, 1, 8), + (2560, 17, 136), + (7168, 64, 512), + ], +) + + +def _inputs(k, num_src_rows, num_dst_rows, seed): + torch.manual_seed(seed) + q = torch.randint(0, 256, (num_src_rows, k), dtype=torch.uint8, device="cuda").view( + torch.float8_e4m3fn + ) + scale = torch.randn( + (num_src_rows, k // GROUP_SIZE), dtype=torch.float32, device="cuda" + ) + # Duplicate source rows are the normal case here: a token is replicated once + # per expert it routes to. + dst2src = torch.randint( + 0, num_src_rows, (num_dst_rows,), dtype=torch.int32, device="cuda" + ) + return q, scale, dst2src + + +def _assert_same_bytes(got, ref, what): + assert torch.equal( + got.view(torch.int8), ref.view(torch.int8) + ), f"{what} bytes differ" + + +@pytest.mark.parametrize("k,num_src_rows,num_dst_rows", CASES) +def test_matches_torch_gather(k, num_src_rows, num_dst_rows): + q, scale, dst2src = _inputs(k, num_src_rows, num_dst_rows, seed=0) + + got_q, got_scale = shuffle_rows_with_scales(q, scale, dst2src, num_dst_rows) + + idx = dst2src.long() + _assert_same_bytes(got_q, q[idx], "values") + _assert_same_bytes(got_scale, scale[idx], "scales") + + +CUDA_XCHECK_CASES = [ + c for c in CASES if (c[0] // GROUP_SIZE) % CUDA_XCHECK_SCALE_ALIGN == 0 +] +# An empty parametrize list collects zero tests and reports success, which would +# retire the drop-in check without saying so. Fail at collection instead. +assert CUDA_XCHECK_CASES, "no case survives the CUDA cross-check shape filter" + + +@pytest.mark.parametrize("k,num_src_rows,num_dst_rows", CUDA_XCHECK_CASES) +def test_matches_shuffle_rows_pair(k, num_src_rows, num_dst_rows): + """Drop-in equivalence with the two launches this replaces.""" + from sgl_kernel import shuffle_rows + + q, scale, dst2src = _inputs(k, num_src_rows, num_dst_rows, seed=1) + + got_q, got_scale = shuffle_rows_with_scales(q, scale, dst2src, num_dst_rows) + + _assert_same_bytes(got_q, shuffle_rows(q, dst2src, (num_dst_rows, k)), "values") + _assert_same_bytes( + got_scale, + shuffle_rows(scale, dst2src, (num_dst_rows, k // GROUP_SIZE)), + "scales", + ) + + +def test_empty_destination_does_not_launch(): + """Zero rows takes the short circuit instead of a zero-sized grid.""" + q, scale, dst2src = _inputs(512, 4, 0, seed=2) + + got_q, got_scale = shuffle_rows_with_scales(q, scale, dst2src, 0) + + assert got_q.shape == (0, 512) + assert got_scale.shape == (0, 512 // GROUP_SIZE) + assert got_q.dtype == q.dtype and got_scale.dtype == scale.dtype + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__]))