From b86c6491fa7080ccea00630935c9c88f9407b847 Mon Sep 17 00:00:00 2001 From: Mohammad Miadh Angkad <176301910+mmangkad@users.noreply.github.com> Date: Mon, 16 Feb 2026 19:47:09 +0800 Subject: [PATCH] [Perf] ~9.5x faster Blackwell MXFP4 MoE weight loading (#18858) --- .../sglang/srt/layers/quantization/mxfp4.py | 114 ++++++++++++++---- 1 file changed, 92 insertions(+), 22 deletions(-) diff --git a/python/sglang/srt/layers/quantization/mxfp4.py b/python/sglang/srt/layers/quantization/mxfp4.py index 3690b4d59..46c51cae0 100644 --- a/python/sglang/srt/layers/quantization/mxfp4.py +++ b/python/sglang/srt/layers/quantization/mxfp4.py @@ -16,7 +16,6 @@ from __future__ import annotations -import logging from typing import TYPE_CHECKING, List, Optional import torch @@ -45,7 +44,6 @@ from sglang.srt.utils import ( is_sm90_supported, is_sm100_supported, is_triton_kernels_available, - log_info_on_rank0, mxfp_supported, next_power_of_2, round_up, @@ -62,12 +60,50 @@ has_triton_kernels = is_triton_kernels_available() if is_flashinfer_available(): from flashinfer import ( mxfp8_quantize, - shuffle_matrix_a, - shuffle_matrix_sf_a, + nvfp4_block_scale_interleave, trtllm_fp4_block_scale_moe, ) + from flashinfer.fused_moe.core import get_w2_permute_indices_with_cache + +_flashinfer_mxfp4_permute_indices_cache: dict[torch.Size, torch.Tensor] = {} +_flashinfer_mxfp4_permute_indices_device_cache: dict[ + tuple[tuple[int, ...], int, int, str, int], torch.Tensor +] = {} + + +def _get_flashinfer_mxfp4_device_permute_indices( + x: torch.Tensor, + epilogue_tile_m: int, + num_elts_per_sf: Optional[int] = None, +) -> torch.Tensor: + extra_args = {} if num_elts_per_sf is None else {"num_elts_per_sf": num_elts_per_sf} + permute_indices = get_w2_permute_indices_with_cache( + _flashinfer_mxfp4_permute_indices_cache, + x, + epilogue_tile_m, + **extra_args, + ) + + device_index = -1 if x.device.index is None else x.device.index + num_elts_per_sf_key = -1 if num_elts_per_sf is None else num_elts_per_sf + cache_key = ( + tuple(x.shape), + epilogue_tile_m, + num_elts_per_sf_key, + x.device.type, + device_index, + ) + cached_device_indices = _flashinfer_mxfp4_permute_indices_device_cache.get( + cache_key + ) + if cached_device_indices is None: + cached_device_indices = permute_indices.to(x.device) + _flashinfer_mxfp4_permute_indices_device_cache[cache_key] = ( + cached_device_indices + ) + + return cached_device_indices -logger = logging.getLogger(__name__) if TYPE_CHECKING: from sglang.srt.layers.moe.token_dispatcher import ( @@ -391,10 +427,6 @@ class Mxfp4MoEMethod(FusedMoEMethodBase): def process_weights_after_loading(self, layer): if self.use_flashinfer: - log_info_on_rank0( - logger, - f"Shuffling MoE weights for FlashInfer MXFP4 moe kernel (layer: {self.prefix}), it might take a while...", - ) # TODO: these values are hardcoded for now, we need to get them from the model layer.gemm1_alpha = Parameter( torch.tensor([1.702] * self.num_experts, dtype=torch.float32).cuda(), @@ -486,31 +518,69 @@ class Mxfp4MoEMethod(FusedMoEMethodBase): gemm1_bias_shuffled = [] gemm2_bias_shuffled = [] epilogue_tile_m = 128 # FIXME: this depends on the kernel internals + w13_weight_permute_indices = _get_flashinfer_mxfp4_device_permute_indices( + w13_weight[0].view(torch.uint8), + epilogue_tile_m, + ) + w13_scale_permute_indices = _get_flashinfer_mxfp4_device_permute_indices( + w13_weight_scale[0].view(torch.uint8), + epilogue_tile_m, + num_elts_per_sf=16, + ) + w13_bias_permute_indices = _get_flashinfer_mxfp4_device_permute_indices( + w13_bias[0].reshape(-1, 1), + epilogue_tile_m, + ) + + w2_weight_permute_indices = _get_flashinfer_mxfp4_device_permute_indices( + w2_weight[0].view(torch.uint8), + epilogue_tile_m, + ) + w2_scale_permute_indices = _get_flashinfer_mxfp4_device_permute_indices( + w2_weight_scale[0].view(torch.uint8), + epilogue_tile_m, + num_elts_per_sf=16, + ) + w2_bias_permute_indices = _get_flashinfer_mxfp4_device_permute_indices( + w2_bias[0].reshape(-1, 1), + epilogue_tile_m, + ) + for i in range(self.num_experts): gemm1_weights_mxfp4_shuffled.append( - shuffle_matrix_a(w13_weight[i].view(torch.uint8), epilogue_tile_m) + w13_weight[i] + .view(torch.uint8)[w13_weight_permute_indices] + .contiguous() ) + gemm1_scales_mxfp4_shuffled.append( - shuffle_matrix_sf_a( - w13_weight_scale[i].view(torch.uint8), epilogue_tile_m - ) - ) - gemm1_bias_shuffled.append( - shuffle_matrix_a( - w13_bias[i].clone().reshape(-1, 1), epilogue_tile_m + nvfp4_block_scale_interleave( + w13_weight_scale[i] + .view(torch.uint8)[w13_scale_permute_indices] + .contiguous() ) ) - gemm2_weights_mxfp4_shuffled.append( - shuffle_matrix_a(w2_weight[i].view(torch.uint8), epilogue_tile_m) + gemm1_bias_shuffled.append( + w13_bias[i].reshape(-1, 1)[w13_bias_permute_indices].contiguous() ) + + gemm2_weights_mxfp4_shuffled.append( + w2_weight[i] + .view(torch.uint8)[w2_weight_permute_indices] + .contiguous() + ) + gemm2_scales_mxfp4_shuffled.append( - shuffle_matrix_sf_a( - w2_weight_scale[i].view(torch.uint8), epilogue_tile_m + nvfp4_block_scale_interleave( + w2_weight_scale[i] + .view(torch.uint8)[w2_scale_permute_indices] + .contiguous() ) ) + gemm2_bias_shuffled.append( - shuffle_matrix_a(w2_bias[i].clone().reshape(-1, 1), epilogue_tile_m) + w2_bias[i].reshape(-1, 1)[w2_bias_permute_indices].contiguous() ) w13_weight = torch.stack(gemm1_weights_mxfp4_shuffled)