From b3ee679467fbc0e423d893608fe5753f44e9dc67 Mon Sep 17 00:00:00 2001 From: Brayden Zhong Date: Fri, 7 Aug 2026 14:30:43 -0700 Subject: [PATCH] [MXFP8] Use FlashInfer CUTLASS for dense GEMM on SM120, delete Triton path (#33208) Co-authored-by: Brayden Zhong --- .../kernels/ops/quantization/__init__.py | 1 - .../kernels/ops/quantization/fp8_kernel.py | 210 ----------- python/sglang/srt/layers/quantization/fp8.py | 57 +-- .../srt/layers/quantization/fp8_utils.py | 356 ++++++------------ test/manual/quant/test_block_fp8.py | 36 +- .../gemm/test_minimax_fused_qkv_index_gemm.py | 6 +- 6 files changed, 171 insertions(+), 495 deletions(-) diff --git a/python/sglang/kernels/ops/quantization/__init__.py b/python/sglang/kernels/ops/quantization/__init__.py index a6f9ede74..2b389d604 100644 --- a/python/sglang/kernels/ops/quantization/__init__.py +++ b/python/sglang/kernels/ops/quantization/__init__.py @@ -172,7 +172,6 @@ _TRITON_KERNELS = [ ("fp8_kernel", "sglang_per_token_quant_fp8"), ("fp8_kernel", "static_quant_fp8"), ("fp8_kernel", "w8a8_block_fp8_matmul"), - ("fp8_kernel", "mxfp8_block_scaled_matmul_triton"), ("fp8_kernel", "per_tensor_quant_mla_fp8"), ("fp8_kernel", "per_token_group_quant_mla_deep_gemm_masked_fp8"), ("fp8_kernel", "per_token_group_quant_fp8_hopper_moe_mn_major"), diff --git a/python/sglang/kernels/ops/quantization/fp8_kernel.py b/python/sglang/kernels/ops/quantization/fp8_kernel.py index 1ab334b9e..b0b8dec6b 100644 --- a/python/sglang/kernels/ops/quantization/fp8_kernel.py +++ b/python/sglang/kernels/ops/quantization/fp8_kernel.py @@ -23,11 +23,6 @@ import torch import triton import triton.language as tl -try: - from triton.tools.tensor_descriptor import TensorDescriptor -except: - pass - from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.kernels.ops.quantization.fp8_utils import fp8_dtype_to_triton from sglang.srt.layers import deep_gemm_wrapper @@ -40,8 +35,6 @@ from sglang.srt.utils import ( is_cuda, is_hip, is_musa, - is_sm100_supported, - is_sm120_supported, log_info_on_rank0, ) from sglang.srt.utils.custom_op import register_custom_op @@ -51,8 +44,6 @@ _is_hip = is_hip() _is_cuda = is_cuda() _is_cpu = is_cpu() _is_musa = is_musa() -_is_sm100_supported = is_sm100_supported() -_is_sm120_supported = is_sm120_supported() _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip if _is_cuda or _is_musa: @@ -1474,207 +1465,6 @@ def w8a8_block_fp8_matmul( ) -# Copied and adapted from https://github.com/triton-lang/triton/blob/main/python/tutorials/10-block-scaled-matmul.py -@triton.jit -def _mxfp8_block_scaled_matmul_kernel( # - a_desc, # - a_scale_desc, # - b_desc, # - b_scale_desc, # - c_desc, # - M: tl.constexpr, # - N: tl.constexpr, # - K: tl.constexpr, # - output_type: tl.constexpr, # - BLOCK_M: tl.constexpr, # - BLOCK_N: tl.constexpr, # - BLOCK_K: tl.constexpr, # - rep_m: tl.constexpr, # - rep_n: tl.constexpr, # - rep_k: tl.constexpr, # - NUM_STAGES: tl.constexpr, # -): # - if output_type == 0: - output_dtype = tl.float32 - elif output_type == 1: - output_dtype = tl.float16 - elif output_type == 2: - output_dtype = tl.bfloat16 - - pid = tl.program_id(axis=0) - num_pid_m = tl.cdiv(M, BLOCK_M) - pid_m = pid % num_pid_m - pid_n = pid // num_pid_m - offs_am = pid_m * BLOCK_M - offs_bn = pid_n * BLOCK_N - offs_k_a = 0 - offs_k_b = 0 - offs_scale_m = pid_m * rep_m - offs_scale_n = pid_n * rep_n - offs_scale_k = 0 - - VEC_SIZE: tl.constexpr = 32 - - accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) - for k in tl.range(0, tl.cdiv(K, BLOCK_K), num_stages=NUM_STAGES): - a = a_desc.load([offs_am, offs_k_a]) - b = b_desc.load([offs_bn, offs_k_b]) - scale_a = a_scale_desc.load([0, offs_scale_m, offs_scale_k, 0, 0]) - scale_b = b_scale_desc.load([0, offs_scale_n, offs_scale_k, 0, 0]) - - scale_a = ( - scale_a.reshape(rep_m, rep_k, 32, 4, 4) - .trans(0, 3, 2, 1, 4) - .reshape(BLOCK_M, BLOCK_K // VEC_SIZE) - ) - scale_b = ( - scale_b.reshape(rep_n, rep_k, 32, 4, 4) - .trans(0, 3, 2, 1, 4) - .reshape(BLOCK_N, BLOCK_K // VEC_SIZE) - ) - - accumulator = tl.dot_scaled( - a, scale_a, "e4m3", b.T, scale_b, "e4m3", accumulator - ) - - offs_k_a += BLOCK_K - offs_k_b += BLOCK_K - offs_scale_k += rep_k - - c_desc.store([offs_am, offs_bn], accumulator.to(output_dtype)) - - -# Copied and adapted from https://github.com/triton-lang/triton/blob/main/python/tutorials/10-block-scaled-matmul.py -def mxfp8_block_scaled_matmul_triton( - a: torch.Tensor, - a_scale: torch.Tensor, - b: torch.Tensor, - b_scale: torch.Tensor, - output_dtype: torch.dtype, - *, - block_m: int = 128, - block_n: int = 256, - block_k: int = 128, - num_stages: Optional[int] = None, -) -> torch.Tensor: - """Block-scaled matmul for MXFP8 using Triton dot_scaled. - - Args: - num_stages: Number of pipeline stages. If None, auto-selects based on GPU: - SM120: 1, SM100: 4. - """ - if num_stages is None: - num_stages = 1 if _is_sm120_supported else (4 if _is_sm100_supported else 1) - M, K = a.shape - N, K_b = b.shape - assert K == K_b - - if output_dtype == torch.float32: - output_type = 0 - elif output_dtype == torch.float16: - output_type = 1 - elif output_dtype == torch.bfloat16: - output_type = 2 - else: - raise ValueError(f"Unsupported output dtype: {output_dtype}") - - rep_m = block_m // 128 - rep_n = block_n // 128 - rep_k = block_k // 32 // 4 - - a_desc = TensorDescriptor.from_tensor(a, [block_m, block_k]) - b_desc = TensorDescriptor.from_tensor(b, [block_n, block_k]) - - scale_block_shape = [1, rep_m, rep_k, 2, 256] - a_scale_desc = TensorDescriptor.from_tensor(a_scale, block_shape=scale_block_shape) - scale_block_shape = [1, rep_n, rep_k, 2, 256] - b_scale_desc = TensorDescriptor.from_tensor(b_scale, block_shape=scale_block_shape) - - output = torch.empty((M, N), dtype=output_dtype, device=a.device) - c_desc = TensorDescriptor.from_tensor(output, [block_m, block_n]) - - grid = (triton.cdiv(M, block_m) * triton.cdiv(N, block_n), 1) - _mxfp8_block_scaled_matmul_kernel[grid]( - a_desc, - a_scale_desc, - b_desc, - b_scale_desc, - c_desc, - M, - N, - K, - output_type, - block_m, - block_n, - block_k, - rep_m, - rep_n, - rep_k, - num_stages, - ) - return output - - -@triton.jit -def _pack_mxfp8_scales_kernel( - scale_ptr, - out_ptr, - M: tl.constexpr, - K_GROUPS: tl.constexpr, - SCALE_K: tl.constexpr, - TOTAL: tl.constexpr, - BLOCK: tl.constexpr, -): - offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) - mask = offs < TOTAL - - idx256 = offs % 256 - tmp = offs // 256 - two = tmp % 2 - tmp = tmp // 2 - scale_k = tmp % SCALE_K - scale_m = tmp // SCALE_K - - within = two * 256 + idx256 - row_inner_32 = within // 16 - rem = within - row_inner_32 * 16 - row_outer_4 = rem // 4 - k_inner_4 = rem - row_outer_4 * 4 - - row = scale_m * 128 + row_outer_4 * 32 + row_inner_32 - col = scale_k * 4 + k_inner_4 - value = tl.load(scale_ptr + row * K_GROUPS + col, mask & (row < M), other=127) - tl.store(out_ptr + offs, value, mask) - - -def pack_mxfp8_scales_triton(scale_u8: torch.Tensor) -> torch.Tensor: - assert scale_u8.dim() == 2, f"Expected 2D scale tensor, got {scale_u8.dim()}D" - scale_u8 = scale_u8.contiguous() - m, k_groups = scale_u8.shape - assert ( - k_groups % 4 == 0 - ), f"{k_groups=} must be divisible by 4 (K must be multiple of 128)" - - scale_m = triton.cdiv(m, 128) - scale_k = k_groups // 4 - out = torch.empty( - (1, scale_m, scale_k, 2, 256), dtype=scale_u8.dtype, device=scale_u8.device - ) - total = out.numel() - block = 1024 - grid = (triton.cdiv(total, block),) - _pack_mxfp8_scales_kernel[grid]( - scale_u8, - out, - m, - k_groups, - scale_k, - total, - BLOCK=block, - ) - return out - - @triton.jit def _per_tensor_quant_mla_fp8_stage1( x_ptr, diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py index 26c5150ee..acc40d05c 100644 --- a/python/sglang/srt/layers/quantization/fp8.py +++ b/python/sglang/srt/layers/quantization/fp8.py @@ -60,11 +60,11 @@ from sglang.srt.layers.quantization.fp8_utils import ( deepgemm_w8a8_block_fp8_linear_with_fallback, dispatch_w8a8_block_fp8_linear, dispatch_w8a8_mxfp8_linear, - get_fp8_gemm_runner_backend, input_to_float8, mxfp8_group_quantize, normalize_e4m3fn_to_e4m3fnuz, requant_block_scale_ue8m0_for_deepgemm, + resolve_mxfp8_dense_gemm_backend, ) from sglang.srt.layers.quantization.kv_cache import BaseKVCacheMethod from sglang.srt.layers.quantization.marlin_utils_fp8 import prepare_fp8_layer_for_marlin @@ -84,8 +84,10 @@ from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, + is_blackwell_supported, is_cpu, is_cuda, + is_flashinfer_available, is_gfx95_supported, is_hip, is_musa, @@ -470,7 +472,9 @@ class Fp8LinearMethod(LinearMethodBase): self.weight_block_size = self.quant_config.weight_block_size self.w8a8_block_fp8_linear = None self.w8a8_mxfp8_linear = None + self.mxfp8_dense_backend = None if self.use_mxfp8 and not self.convert_mxfp8_to_block: + self.mxfp8_dense_backend = resolve_mxfp8_dense_gemm_backend() self.w8a8_mxfp8_linear = dispatch_w8a8_mxfp8_linear() else: self.w8a8_block_fp8_linear = dispatch_w8a8_block_fp8_linear() @@ -730,7 +734,7 @@ class Fp8LinearMethod(LinearMethodBase): if not self.use_mxfp8: return - backend = get_fp8_gemm_runner_backend() + backend = self.mxfp8_dense_backend if backend.is_flashinfer_trtllm(): from flashinfer import shuffle_matrix_a, shuffle_matrix_sf_a @@ -783,13 +787,28 @@ class Fp8LinearMethod(LinearMethodBase): "weight_scale_inv_swizzled", block_scale_interleave(scale_u8.contiguous()).contiguous(), ) - elif get_fp8_gemm_runner_backend().is_deep_gemm(): + elif backend.is_deep_gemm(): from sglang.srt.layers.deep_gemm_wrapper.configurer import ( DEEPGEMM_SCALE_UE8M0, ) n, k = layer.weight.shape scale_u8 = layer.weight_scale_inv.data + layer.weight_scale_inv_swizzled = None + if n % 64 != 0 or k % 128 != 0: + if not (is_blackwell_supported() and is_flashinfer_available()): + raise RuntimeError( + f"--fp8-gemm-backend=deep_gemm cannot serve MXFP8 weight shape " + f"({n}, {k}) (needs N % 64 == 0 and K % 128 == 0), and this " + "device has no FlashInfer MXFP8 fallback kernel." + ) + from flashinfer import block_scale_interleave + + copy_or_rebind_param( + layer, + "weight_scale_inv_swizzled", + block_scale_interleave(scale_u8.contiguous()).contiguous(), + ) scale_fp32 = ( (scale_u8.contiguous().view(-1).to(torch.int32) << 23) .view(torch.float32) @@ -807,9 +826,6 @@ class Fp8LinearMethod(LinearMethodBase): else: scale_packed = scale_fp32 copy_or_rebind_param(layer, "weight_scale_inv_deepgemm", scale_packed) - else: - # Triton path consumes canonical 2D UE8M0 uint8 scales directly. - return def _quantize_mxfp8_weights(self, layer: Module) -> None: weight = layer.weight.data @@ -959,32 +975,15 @@ class Fp8LinearMethod(LinearMethodBase): ) if self.use_mxfp8: - backend = get_fp8_gemm_runner_backend() + backend = self.mxfp8_dense_backend + extra_kwargs = {} if backend.is_flashinfer_cutlass(): weight_scale = layer.weight_scale_inv_swizzled elif backend.is_flashinfer_trtllm(): weight_scale = layer.weight_scale_inv_shuffled - elif get_fp8_gemm_runner_backend().is_deep_gemm(): - weight_scale = getattr( - layer, "weight_scale_inv_deepgemm", layer.weight_scale_inv - ) - if isinstance(x, tuple): - return self.w8a8_mxfp8_linear( - input=x[0], - weight=layer.weight, - weight_scale=weight_scale, - input_scale=x[1], - bias=bias, - weight_scale_fallback=layer.weight_scale_inv, - ) - return self.w8a8_mxfp8_linear( - input=x, - weight=layer.weight, - weight_scale=weight_scale, - input_scale=None, - bias=bias, - weight_scale_fallback=layer.weight_scale_inv, - ) + elif backend.is_deep_gemm(): + weight_scale = layer.weight_scale_inv_deepgemm + extra_kwargs["weight_scale_swizzled"] = layer.weight_scale_inv_swizzled else: weight_scale = layer.weight_scale_inv if isinstance(x, tuple): @@ -994,6 +993,7 @@ class Fp8LinearMethod(LinearMethodBase): weight_scale=weight_scale, input_scale=x[1], bias=bias, + **extra_kwargs, ) return self.w8a8_mxfp8_linear( input=x, @@ -1001,6 +1001,7 @@ class Fp8LinearMethod(LinearMethodBase): weight_scale=weight_scale, input_scale=None, bias=bias, + **extra_kwargs, ) if self.block_quant: diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index da675c84b..3a31c3905 100755 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from enum import Enum -from functools import lru_cache +from functools import lru_cache, partial from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Union import torch @@ -24,8 +24,6 @@ from sglang.kernels.ops.quantization.fp8_kernel import ( fp8_max, fp8_min, is_fp8_fnuz, - mxfp8_block_scaled_matmul_triton, - pack_mxfp8_scales_triton, per_token_group_quant_fp8, scaled_fp8_quant, sglang_per_token_quant_fp8, @@ -60,7 +58,6 @@ _is_hip = is_hip() _is_cuda = is_cuda() _is_fp8_fnuz = is_fp8_fnuz() _is_sm100_supported = is_sm100_supported() -_is_sm120_supported = is_sm120_supported() _is_gfx95_supported = is_gfx95_supported() _is_musa = is_musa() @@ -308,6 +305,32 @@ class Fp8GemmRunnerBackend(Enum): return self == Fp8GemmRunnerBackend.AITER +class Mxfp8DenseGemmBackend(Enum): + """Enum for MXFP8 dense linear backend selection, resolved separately from + `Fp8GemmRunnerBackend`.""" + + FLASHINFER_CUTLASS = "flashinfer_cutlass" + FLASHINFER_TRTLLM = "flashinfer_trtllm" + DEEP_GEMM = "deep_gemm" + GFX95_DOT_SCALED = "gfx95_dot_scaled" + UNSUPPORTED = "unsupported" + + def is_flashinfer_cutlass(self) -> bool: + return self == Mxfp8DenseGemmBackend.FLASHINFER_CUTLASS + + def is_flashinfer_trtllm(self) -> bool: + return self == Mxfp8DenseGemmBackend.FLASHINFER_TRTLLM + + def is_deep_gemm(self) -> bool: + return self == Mxfp8DenseGemmBackend.DEEP_GEMM + + def is_gfx95_dot_scaled(self) -> bool: + return self == Mxfp8DenseGemmBackend.GFX95_DOT_SCALED + + def is_unsupported(self) -> bool: + return self == Mxfp8DenseGemmBackend.UNSUPPORTED + + FP8_GEMM_RUNNER_BACKEND: Fp8GemmRunnerBackend | None = None @@ -495,21 +518,64 @@ def dispatch_w8a8_block_fp8_linear() -> Callable: return _dispatch_auto_backend() -def dispatch_w8a8_mxfp8_linear() -> Callable: +def resolve_mxfp8_dense_gemm_backend() -> Mxfp8DenseGemmBackend: + """Pick the MXFP8 dense linear backend, honoring `--fp8-gemm-backend` only when it + names a backend that owns an MXFP8 dense kernel.""" backend = get_fp8_gemm_runner_backend() + + if backend.is_flashinfer_trtllm(): + if not (_is_sm100_supported and is_flashinfer_available()): + raise RuntimeError( + "MXFP8 dense GEMM requested via --fp8-gemm-backend=flashinfer_trtllm, " + "but that kernel requires SM100/SM103 GPUs and FlashInfer." + ) + return Mxfp8DenseGemmBackend.FLASHINFER_TRTLLM + + if backend.is_deep_gemm(): + if not deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM: + raise RuntimeError( + "MXFP8 dense GEMM requested via --fp8-gemm-backend=deep_gemm, but " + "DeepGEMM is not available (package missing or " + "SGLANG_ENABLE_JIT_DEEPGEMM=0)." + ) + return Mxfp8DenseGemmBackend.DEEP_GEMM + + if _is_hip and _is_gfx95_supported: + return Mxfp8DenseGemmBackend.GFX95_DOT_SCALED + + if is_blackwell_supported() and is_flashinfer_available(): + return Mxfp8DenseGemmBackend.FLASHINFER_CUTLASS + + if deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM: + return Mxfp8DenseGemmBackend.DEEP_GEMM + + return Mxfp8DenseGemmBackend.UNSUPPORTED + + +def _unsupported_mxfp8_linear(*args, **kwargs) -> torch.Tensor: + raise RuntimeError( + "No MXFP8 dense GEMM kernel is available on this device. MXFP8 dense linear " + "requires Blackwell (SM100/SM103/SM110/SM120) with FlashInfer, Hopper (SM90) " + "with DeepGEMM, or ROCm gfx95." + ) + + +def dispatch_w8a8_mxfp8_linear() -> Callable: + backend = resolve_mxfp8_dense_gemm_backend() if backend.is_deep_gemm(): return _deepgemm_w8a8_mxfp8_linear_with_fallback - elif backend.is_flashinfer_cutlass() or backend.is_flashinfer_trtllm(): - return flashinfer_mxfp8_blockscaled_linear - elif backend.is_triton(): - return triton_mxfp8_blockscaled_linear - elif _is_hip and _is_gfx95_supported: - from sglang.kernels.ops.quantization.mxfp8_amd_gfx95 import ( - dot_scaled_mxfp8_blockscaled_linear, - ) + elif backend.is_flashinfer_trtllm(): + return partial(flashinfer_mxfp8_blockscaled_linear, backend="trtllm") + elif backend.is_flashinfer_cutlass(): + return partial(flashinfer_mxfp8_blockscaled_linear, backend="cutlass") + elif backend.is_unsupported(): + return _unsupported_mxfp8_linear - return dot_scaled_mxfp8_blockscaled_linear - return triton_mxfp8_blockscaled_linear + from sglang.kernels.ops.quantization.mxfp8_amd_gfx95 import ( + dot_scaled_mxfp8_blockscaled_linear, + ) + + return dot_scaled_mxfp8_blockscaled_linear def _deepgemm_w8a8_mxfp8_linear_with_fallback( @@ -518,7 +584,7 @@ def _deepgemm_w8a8_mxfp8_linear_with_fallback( weight_scale: torch.Tensor, input_scale: Optional[torch.Tensor] = None, bias: Optional[torch.Tensor] = None, - weight_scale_fallback: Optional[torch.Tensor] = None, + weight_scale_swizzled: Optional[torch.Tensor] = None, ) -> torch.Tensor: from sglang.kernels.ops.quantization.fp8_kernel import ( sglang_per_token_group_quant_fp8, @@ -532,8 +598,18 @@ def _deepgemm_w8a8_mxfp8_linear_with_fallback( dtype_supported = output_dtype == torch.bfloat16 if not (shape_supported and dtype_supported): - return triton_mxfp8_blockscaled_linear( - input, weight, weight_scale_fallback, input_scale, bias + if weight_scale_swizzled is None: + raise RuntimeError( + f"DeepGEMM cannot serve this MXFP8 GEMM ({shape_supported=}, " + f"{dtype_supported=}) and no FlashInfer fallback scale was prepared " + "for this layer. Re-run with --fp8-gemm-backend=flashinfer_cutlass." + ) + return flashinfer_mxfp8_blockscaled_linear( + input=input, + weight=weight, + weight_scale=weight_scale_swizzled, + input_scale=input_scale, + bias=bias, ) input_2d = input.view(-1, input.shape[-1]) @@ -1132,197 +1208,6 @@ def mxfp8_group_quantize(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: return q_input.contiguous(), scale_u8.contiguous() -def _pack_mxfp8_scales(scale_u8: torch.Tensor) -> torch.Tensor: - if ( - _is_hip - and _is_gfx95_supported - and scale_u8.is_cuda - and scale_u8.shape[0] % 128 != 0 - ): - return pack_mxfp8_scales_triton(scale_u8) - - # Pack (M, K//32) UE8M0 scales into the layout expected by tl.dot_scaled. - assert scale_u8.dim() == 2, f"Expected 2D scale tensor, got {scale_u8.dim()}D" - scale_u8 = scale_u8.contiguous() - m, k_groups = scale_u8.shape - assert ( - k_groups % 4 == 0 - ), f"{k_groups=} must be divisible by 4 (K must be multiple of 128)" - - scale_m = ceil_div(m, 128) - if m % 128 != 0: - pad_rows = scale_m * 128 - m - pad = torch.full( - (pad_rows, k_groups), - 127, - dtype=scale_u8.dtype, - device=scale_u8.device, - ) - scale_u8 = torch.cat([scale_u8, pad], dim=0) - - scale_k = k_groups // 4 - scale_u8 = scale_u8.view(scale_m, 128, scale_k, 4) - scale_u8 = scale_u8.view(scale_m, 4, 32, scale_k, 4) - packed = scale_u8.permute(0, 3, 2, 1, 4).contiguous() - return packed.view(1, scale_m, scale_k, 2, 256) - - -@register_custom_op( - op_name="triton_mxfp8_block_scaled_matmul", - mutates_args=[], - fake_impl=lambda a, a_scale, b, b_scale, output_dtype, block_m=128, block_n=256, block_k=128, num_stages=None: ( # noqa: E501 - a.new_empty((a.shape[0], b.shape[0]), dtype=output_dtype) - ), -) -def triton_mxfp8_block_scaled_matmul( - a: torch.Tensor, - a_scale: torch.Tensor, - b: torch.Tensor, - b_scale: torch.Tensor, - output_dtype: torch.dtype, - *, - block_m: int = 128, - block_n: int = 256, - block_k: int = 128, - num_stages: Optional[int] = None, -) -> torch.Tensor: - """Opaque custom op wrapper to prevent Dynamo tracing Triton grid math.""" - return mxfp8_block_scaled_matmul_triton( - a, - a_scale, - b, - b_scale, - output_dtype=output_dtype, - block_m=block_m, - block_n=block_n, - block_k=block_k, - num_stages=num_stages, - ) - - -def _raw_triton_mxfp8_blockscaled_linear( - input: torch.Tensor, - weight: torch.Tensor, - weight_scale: torch.Tensor, - input_scale: Optional[torch.Tensor] = None, - bias: Optional[torch.Tensor] = None, - output_dtype: Optional[torch.dtype] = None, -) -> torch.Tensor: - if not ( - (_is_cuda and (_is_sm100_supported or _is_sm120_supported)) - or (_is_hip and _is_gfx95_supported) - ): - raise RuntimeError( - "MXFP8 dense linear requires Blackwell GPUs (SM100/SM120) or ROCm gfx95." - ) - - input_2d = input.view(-1, input.shape[-1]).contiguous() - output_shape = [*input.shape[:-1], weight.shape[0]] - - block_m = 128 - block_n = 256 if weight.shape[0] % 256 == 0 else 128 - block_k = 128 - - m, k = input_2d.shape - n, k_w = weight.shape - assert k == k_w, f"{k=} does not match {k_w=}" - assert k % 128 == 0, f"{k=} must be divisible by 128 for MXFP8" - assert n % block_n == 0, f"{n=} must be divisible by {block_n}" - assert weight.dtype == torch.float8_e4m3fn, "MXFP8 weight must be FP8 E4M3." - assert weight_scale.dtype == torch.uint8, "MXFP8 weight_scale must be UE8M0 uint8." - assert weight_scale.dim() in ( - 2, - 5, - ), ( - "MXFP8 weight_scale must be canonical 2D or packed 5D, " - f"got {weight_scale.dim()}D." - ) - - if input_scale is None: - q_input, x_scale_u8 = mxfp8_group_quantize(input_2d) - else: - q_input = input_2d - x_scale_u8 = input_scale - assert x_scale_u8.dtype == torch.uint8, "MXFP8 input_scale must be UE8M0 uint8." - assert x_scale_u8.shape == (m, k // 32) - - if output_dtype is None: - if input_2d.dtype in (torch.float16, torch.bfloat16, torch.float32): - output_dtype = input_2d.dtype - else: - output_dtype = torch.bfloat16 - - if m % block_m != 0: - pad_rows = ceil_div(m, block_m) * block_m - m - q_input = torch.cat( - [ - q_input, - torch.zeros((pad_rows, k), device=q_input.device, dtype=q_input.dtype), - ], - dim=0, - ) - pad_scale = torch.full( - (pad_rows, k // 32), - 127, - device=x_scale_u8.device, - dtype=x_scale_u8.dtype, - ) - x_scale_u8 = torch.cat([x_scale_u8, pad_scale], dim=0) - - a_scale_packed = _pack_mxfp8_scales(x_scale_u8) - b_scale_packed = ( - weight_scale.contiguous() - if weight_scale.dim() == 5 - else _pack_mxfp8_scales(weight_scale) - ) - - num_stages = 1 if _is_sm120_supported else (4 if _is_sm100_supported else 1) - output = triton_mxfp8_block_scaled_matmul( - q_input, - a_scale_packed, - weight.contiguous(), - b_scale_packed, - output_dtype=output_dtype, - block_m=block_m, - block_n=block_n, - block_k=block_k, - num_stages=num_stages, - ) - output = output[:m, :] - if bias is not None: - output += bias - return output.to(dtype=output_dtype).view(*output_shape) - - -@register_custom_op( - op_name="triton_mxfp8_blockscaled_linear", - mutates_args=[], - fake_impl=lambda input, weight, weight_scale, input_scale=None, bias=None, output_dtype=None: ( - input.new_empty( - (*input.shape[:-1], weight.shape[0]), - dtype=(output_dtype if output_dtype is not None else input.dtype), - ) - ), -) -def triton_mxfp8_blockscaled_linear( - input: torch.Tensor, - weight: torch.Tensor, - weight_scale: torch.Tensor, - input_scale: Optional[torch.Tensor] = None, - bias: Optional[torch.Tensor] = None, - output_dtype: Optional[torch.dtype] = None, -) -> torch.Tensor: - """Opaque custom-op wrapper to prevent Dynamo guards on MXFP8 padding branches.""" - return _raw_triton_mxfp8_blockscaled_linear( - input=input, - weight=weight, - weight_scale=weight_scale, - input_scale=input_scale, - bias=bias, - output_dtype=output_dtype, - ) - - def flashinfer_mxfp8_blockscaled_linear( input: torch.Tensor, weight: torch.Tensor, @@ -1330,13 +1215,15 @@ def flashinfer_mxfp8_blockscaled_linear( input_scale: Optional[torch.Tensor] = None, bias: Optional[torch.Tensor] = None, output_dtype: Optional[torch.dtype] = None, + backend: str = "cutlass", ) -> torch.Tensor: - """MXFP8 dense linear via FlashInfer mm_mxfp8.""" - input_2d = input.view(-1, input.shape[-1]).contiguous() + """MXFP8 dense linear via FlashInfer mm_mxfp8. `weight_scale` must be the layout + the backend expects, prepared at load time.""" + input_2d = input.view(-1, input.shape[-1]) output_shape = [*input.shape[:-1], weight.shape[0]] - m, k = input_2d.shape - n, k_w = weight.shape + k = input_2d.shape[1] + k_w = weight.shape[1] if k != k_w: raise ValueError(f"Input K={k} does not match weight K={k_w}.") if k % 32 != 0: @@ -1350,7 +1237,7 @@ def flashinfer_mxfp8_blockscaled_linear( ) else: q_input = input_2d - x_scale_u8 = input_scale.contiguous() + x_scale_u8 = input_scale if output_dtype is None: if input_2d.dtype in (torch.float16, torch.bfloat16, torch.float32): @@ -1358,35 +1245,20 @@ def flashinfer_mxfp8_blockscaled_linear( else: output_dtype = torch.bfloat16 - # Ensure transposed tensors are contiguous for FlashInfer's internal runner. - weight_t = weight.contiguous().t() + if backend == "trtllm": + weight_scale_t = weight_scale.view(-1) + else: + weight_scale_t = weight_scale.t() if weight_scale.ndim == 2 else weight_scale - if get_fp8_gemm_runner_backend().is_flashinfer_trtllm(): - weight_scale_t = weight_scale.contiguous().view(-1) - output = flashinfer_mm_mxfp8( - q_input, - weight_t, - x_scale_u8, - weight_scale_t, - out_dtype=output_dtype, - use_8x4_sf_layout=False, - backend="trtllm", - ) - elif get_fp8_gemm_runner_backend().is_flashinfer_cutlass(): - weight_scale_t = ( - weight_scale.contiguous().t() - if weight_scale.ndim == 2 - else weight_scale.contiguous() - ) - output = flashinfer_mm_mxfp8( - q_input, - weight_t, - x_scale_u8, - weight_scale_t, - out_dtype=output_dtype, - use_8x4_sf_layout=False, - backend="cutlass", - ) + output = flashinfer_mm_mxfp8( + q_input, + weight.t(), + x_scale_u8, + weight_scale_t, + out_dtype=output_dtype, + use_8x4_sf_layout=False, + backend=backend, + ) if bias is not None: output += bias diff --git a/test/manual/quant/test_block_fp8.py b/test/manual/quant/test_block_fp8.py index 6d6202dfb..b45eeaa6e 100644 --- a/test/manual/quant/test_block_fp8.py +++ b/test/manual/quant/test_block_fp8.py @@ -17,9 +17,8 @@ from sglang.srt.layers.moe.topk import TopKConfig, select_experts from sglang.srt.layers.quantization.fp8_utils import ( input_to_float8, mxfp8_group_quantize, - triton_mxfp8_blockscaled_linear, ) -from sglang.srt.utils import is_sm100_supported, is_sm120_supported +from sglang.srt.utils import is_blackwell_supported, is_flashinfer_available from sglang.test.test_utils import CustomTestCase _is_cuda = torch.cuda.is_available() and torch.version.cuda @@ -452,11 +451,20 @@ class TestMXFP8DenseLinear(CustomTestCase): def setUpClass(cls): if not torch.cuda.is_available(): raise unittest.SkipTest("CUDA is not available") - if not (is_sm100_supported() or is_sm120_supported()): - raise unittest.SkipTest("MXFP8 requires Blackwell (SM100/SM120)") + if not (is_blackwell_supported() and is_flashinfer_available()): + raise unittest.SkipTest( + "MXFP8 dense linear requires Blackwell + FlashInfer" + ) torch.set_default_device("cuda") def _mxfp8_dense_linear(self, M, NK, dtype, seed): + from flashinfer import block_scale_interleave + + from sglang.srt.layers.quantization.fp8_utils import ( + flashinfer_mxfp8_blockscaled_linear, + flashinfer_mxfp8_quantize, + ) + N, K = NK torch.manual_seed(seed) @@ -465,23 +473,29 @@ class TestMXFP8DenseLinear(CustomTestCase): weight_fp32 = torch.randn((N, K), dtype=torch.float32) / 4 weight_q, weight_scale_u8 = mxfp8_group_quantize(weight_fp32) + weight_scale_swizzled = block_scale_interleave(weight_scale_u8).contiguous() with torch.inference_mode(): - q_input, input_scale_u8 = mxfp8_group_quantize(input_fp16.to(torch.float32)) - a_dq = _mxfp8_group_dequant(q_input, input_scale_u8) + q_input, input_scale_swizzled = flashinfer_mxfp8_quantize( + input_fp16, is_sf_swizzled_layout=True, alignment=32 + ) + _, input_scale_linear = flashinfer_mxfp8_quantize( + input_fp16, is_sf_swizzled_layout=False, alignment=32 + ) + a_dq = _mxfp8_group_dequant(q_input, input_scale_linear.view(M, K // 32)) b_dq = _mxfp8_group_dequant(weight_q, weight_scale_u8) ref_out = torch.matmul(a_dq, b_dq.t()).to(dtype) - out = triton_mxfp8_blockscaled_linear( + out = flashinfer_mxfp8_blockscaled_linear( input=input_fp16, weight=weight_q, - weight_scale=weight_scale_u8, + weight_scale=weight_scale_swizzled, ) - out_prequant = triton_mxfp8_blockscaled_linear( + out_prequant = flashinfer_mxfp8_blockscaled_linear( input=q_input, weight=weight_q, - weight_scale=weight_scale_u8, - input_scale=input_scale_u8, + weight_scale=weight_scale_swizzled, + input_scale=input_scale_swizzled, output_dtype=dtype, ) diff --git a/test/registered/kernels/ops/gemm/test_minimax_fused_qkv_index_gemm.py b/test/registered/kernels/ops/gemm/test_minimax_fused_qkv_index_gemm.py index 4c79af7f4..207f97636 100644 --- a/test/registered/kernels/ops/gemm/test_minimax_fused_qkv_index_gemm.py +++ b/test/registered/kernels/ops/gemm/test_minimax_fused_qkv_index_gemm.py @@ -48,14 +48,14 @@ def test_fused_equals_separate(T, N1, N2, K): x = torch.randn(T, K, dtype=torch.bfloat16, device=dev) s1p, s2p = _pack_weight_scale(s1), _pack_weight_scale(s2) - out1 = mxfp8_linear(x, w1, s1p, weight_scale_fallback=s1) - out2 = mxfp8_linear(x, w2, s2p, weight_scale_fallback=s2) + out1 = mxfp8_linear(x, w1, s1p) + out2 = mxfp8_linear(x, w2, s2p) ref = torch.cat([out1, out2], dim=-1) w = torch.cat([w1, w2], dim=0).contiguous() s = torch.cat([s1, s2], dim=0).contiguous() sp = _pack_weight_scale(s) - fused = mxfp8_linear(x, w, sp, weight_scale_fallback=s) + fused = mxfp8_linear(x, w, sp) assert fused.shape == ref.shape assert torch.equal(