[Perf] ~9.5x faster Blackwell MXFP4 MoE weight loading (#18858)
This commit is contained in:
committed by
GitHub
parent
4f0409f8aa
commit
b86c6491fa
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user