[loader] Reduce transient allocations in NVFP4 MoE setup (#26861)

This commit is contained in:
Yinghai Lu
2026-06-03 21:13:25 -07:00
committed by GitHub
parent 5c8a04ac4e
commit 71c759ebb7
2 changed files with 68 additions and 50 deletions
@@ -1799,10 +1799,15 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
)
layer.register_parameter("w13_weight_scale", w13_weight_scale)
# Only use `swizzle_blockscale` for shapes, not for real content
layer.w13_blockscale_swizzled = Parameter(
swizzle_blockscale(layer.w13_weight_scale), requires_grad=False
)
# TRTLLM replaces blockscale_swizzled with an alias to weight_scale
# during process_weights_after_loading, so skip the expensive
# swizzle+allocate here to avoid GPU memory fragmentation
if self.enable_flashinfer_trtllm_moe:
layer.w13_blockscale_swizzled = None
else:
layer.w13_blockscale_swizzled = Parameter(
swizzle_blockscale(layer.w13_weight_scale), requires_grad=False
)
w2_weight_scale = ModelWeightParameter(
data=torch.empty(
@@ -1817,9 +1822,12 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
)
layer.register_parameter("w2_weight_scale", w2_weight_scale)
layer.w2_blockscale_swizzled = Parameter(
swizzle_blockscale(layer.w2_weight_scale), requires_grad=False
)
if self.enable_flashinfer_trtllm_moe:
layer.w2_blockscale_swizzled = None
else:
layer.w2_blockscale_swizzled = Parameter(
swizzle_blockscale(layer.w2_weight_scale), requires_grad=False
)
from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported
+53 -43
View File
@@ -687,10 +687,31 @@ def prepare_static_weights_for_trtllm_fp4_moe(
num_experts, hidden_size, intermediate_size // 16
) # fp8 scaling factors
gemm1_weights_fp4_shuffled = []
gemm1_scales_fp4_shuffled = []
gemm2_weights_fp4_shuffled = []
gemm2_scales_fp4_shuffled = []
# Pre-allocate output tensors so per-expert shuffles write directly into
# contiguous slices instead of building lists + torch.stack(). This avoids
# O(num_experts) transient GPU allocations whose freed blocks fragment the
# CUDA address space
gemm1_weights_fp4_shuffled = torch.empty_like(gemm1_weights_fp4.view(torch.uint8))
gemm2_weights_fp4_shuffled = torch.empty_like(gemm2_weights_fp4.view(torch.uint8))
# Pre-allocate scale output tensors and a reusable scratch buffer for
# the permuted input to nvfp4_block_scale_interleave.
# nvfp4_block_scale_interleave flattens its input to 1-D, so the
# per-expert output size equals the per-expert input numel.
def _alloc_scale_buffers(scales):
per_expert_shape = scales[0].view(torch.uint8).shape
per_expert_numel = scales[0].numel()
output = scales.new_empty((num_experts, per_expert_numel), dtype=torch.uint8)
scratch = torch.empty(per_expert_shape, dtype=torch.uint8, device=scales.device)
return output, scratch
gemm1_scales_fp4_shuffled, g1s_scratch = _alloc_scale_buffers(
gemm1_scales_linear_fp4
)
gemm2_scales_fp4_shuffled, g2s_scratch = _alloc_scale_buffers(
gemm2_scales_linear_fp4
)
for i in range(num_experts):
permute_indices = _maybe_get_cached_w3_w1_permute_indices(
_cache_permute_indices,
@@ -698,11 +719,9 @@ def prepare_static_weights_for_trtllm_fp4_moe(
epilogue_tile_m,
is_gated_act_gemm=is_gated,
)
gemm1_weights_fp4_shuffled.append(
gemm1_weights_fp4[i]
.view(torch.uint8)[permute_indices.to(gemm1_weights_fp4.device)]
.contiguous()
)
gemm1_weights_fp4_shuffled[i] = gemm1_weights_fp4[i].view(torch.uint8)[
permute_indices.to(gemm1_weights_fp4.device)
]
permute_sf_indices = _maybe_get_cached_w3_w1_permute_indices(
_cache_permute_indices,
@@ -711,26 +730,23 @@ def prepare_static_weights_for_trtllm_fp4_moe(
num_elts_per_sf=16,
is_gated_act_gemm=is_gated,
)
gemm1_scales_fp4_shuffled.append(
nvfp4_block_scale_interleave(
gemm1_scales_linear_fp4[i]
.view(torch.uint8)[
permute_sf_indices.to(gemm1_scales_linear_fp4.device)
]
.contiguous()
)
# Reuse scratch buffer for the permuted scale input
torch.index_select(
gemm1_scales_linear_fp4[i].view(torch.uint8),
0,
permute_sf_indices.to(gemm1_scales_linear_fp4.device),
out=g1s_scratch,
)
gemm1_scales_fp4_shuffled[i] = nvfp4_block_scale_interleave(g1s_scratch)
permute_indices = get_w2_permute_indices_with_cache(
_cache_permute_indices,
gemm2_weights_fp4[i].view(torch.uint8),
epilogue_tile_m,
)
gemm2_weights_fp4_shuffled.append(
gemm2_weights_fp4[i]
.view(torch.uint8)[permute_indices.to(gemm2_weights_fp4.device)]
.contiguous()
)
gemm2_weights_fp4_shuffled[i] = gemm2_weights_fp4[i].view(torch.uint8)[
permute_indices.to(gemm2_weights_fp4.device)
]
permute_sf_indices = get_w2_permute_indices_with_cache(
_cache_permute_indices,
@@ -738,30 +754,24 @@ def prepare_static_weights_for_trtllm_fp4_moe(
epilogue_tile_m,
num_elts_per_sf=16,
)
gemm2_scales_fp4_shuffled.append(
nvfp4_block_scale_interleave(
gemm2_scales_linear_fp4[i]
.view(torch.uint8)[
permute_sf_indices.to(gemm2_scales_linear_fp4.device)
]
.contiguous()
)
torch.index_select(
gemm2_scales_linear_fp4[i].view(torch.uint8),
0,
permute_sf_indices.to(gemm2_scales_linear_fp4.device),
out=g2s_scratch,
)
gemm2_scales_fp4_shuffled[i] = nvfp4_block_scale_interleave(g2s_scratch)
# Stack weights for all experts
gemm1_weights_fp4_shuffled = torch.stack(gemm1_weights_fp4_shuffled)
gemm1_scales_fp4_shuffled = (
torch.stack(gemm1_scales_fp4_shuffled)
.view(torch.float8_e4m3fn)
.reshape(num_experts, gemm1_rows, hidden_size // 16)
)
del g1s_scratch, g2s_scratch
gemm2_weights_fp4_shuffled = torch.stack(gemm2_weights_fp4_shuffled)
gemm2_scales_fp4_shuffled = (
torch.stack(gemm2_scales_fp4_shuffled)
.view(torch.float8_e4m3fn)
.reshape(num_experts, hidden_size, intermediate_size // 16)
)
# Weight outputs stay as uint8 (FP4 packed) — the TRTLLM kernel expects this.
gemm1_scales_fp4_shuffled = gemm1_scales_fp4_shuffled.view(
torch.float8_e4m3fn
).reshape(num_experts, gemm1_rows, hidden_size // 16)
gemm2_scales_fp4_shuffled = gemm2_scales_fp4_shuffled.view(
torch.float8_e4m3fn
).reshape(num_experts, hidden_size, intermediate_size // 16)
return (
gemm1_weights_fp4_shuffled,
gemm1_scales_fp4_shuffled,