[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) layer.register_parameter("w13_weight_scale", w13_weight_scale)
# Only use `swizzle_blockscale` for shapes, not for real content # TRTLLM replaces blockscale_swizzled with an alias to weight_scale
layer.w13_blockscale_swizzled = Parameter( # during process_weights_after_loading, so skip the expensive
swizzle_blockscale(layer.w13_weight_scale), requires_grad=False # 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( w2_weight_scale = ModelWeightParameter(
data=torch.empty( data=torch.empty(
@@ -1817,9 +1822,12 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
) )
layer.register_parameter("w2_weight_scale", w2_weight_scale) layer.register_parameter("w2_weight_scale", w2_weight_scale)
layer.w2_blockscale_swizzled = Parameter( if self.enable_flashinfer_trtllm_moe:
swizzle_blockscale(layer.w2_weight_scale), requires_grad=False 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 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 num_experts, hidden_size, intermediate_size // 16
) # fp8 scaling factors ) # fp8 scaling factors
gemm1_weights_fp4_shuffled = [] # Pre-allocate output tensors so per-expert shuffles write directly into
gemm1_scales_fp4_shuffled = [] # contiguous slices instead of building lists + torch.stack(). This avoids
gemm2_weights_fp4_shuffled = [] # O(num_experts) transient GPU allocations whose freed blocks fragment the
gemm2_scales_fp4_shuffled = [] # 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): for i in range(num_experts):
permute_indices = _maybe_get_cached_w3_w1_permute_indices( permute_indices = _maybe_get_cached_w3_w1_permute_indices(
_cache_permute_indices, _cache_permute_indices,
@@ -698,11 +719,9 @@ def prepare_static_weights_for_trtllm_fp4_moe(
epilogue_tile_m, epilogue_tile_m,
is_gated_act_gemm=is_gated, is_gated_act_gemm=is_gated,
) )
gemm1_weights_fp4_shuffled.append( gemm1_weights_fp4_shuffled[i] = gemm1_weights_fp4[i].view(torch.uint8)[
gemm1_weights_fp4[i] permute_indices.to(gemm1_weights_fp4.device)
.view(torch.uint8)[permute_indices.to(gemm1_weights_fp4.device)] ]
.contiguous()
)
permute_sf_indices = _maybe_get_cached_w3_w1_permute_indices( permute_sf_indices = _maybe_get_cached_w3_w1_permute_indices(
_cache_permute_indices, _cache_permute_indices,
@@ -711,26 +730,23 @@ def prepare_static_weights_for_trtllm_fp4_moe(
num_elts_per_sf=16, num_elts_per_sf=16,
is_gated_act_gemm=is_gated, is_gated_act_gemm=is_gated,
) )
gemm1_scales_fp4_shuffled.append( # Reuse scratch buffer for the permuted scale input
nvfp4_block_scale_interleave( torch.index_select(
gemm1_scales_linear_fp4[i] gemm1_scales_linear_fp4[i].view(torch.uint8),
.view(torch.uint8)[ 0,
permute_sf_indices.to(gemm1_scales_linear_fp4.device) permute_sf_indices.to(gemm1_scales_linear_fp4.device),
] out=g1s_scratch,
.contiguous()
)
) )
gemm1_scales_fp4_shuffled[i] = nvfp4_block_scale_interleave(g1s_scratch)
permute_indices = get_w2_permute_indices_with_cache( permute_indices = get_w2_permute_indices_with_cache(
_cache_permute_indices, _cache_permute_indices,
gemm2_weights_fp4[i].view(torch.uint8), gemm2_weights_fp4[i].view(torch.uint8),
epilogue_tile_m, epilogue_tile_m,
) )
gemm2_weights_fp4_shuffled.append( gemm2_weights_fp4_shuffled[i] = gemm2_weights_fp4[i].view(torch.uint8)[
gemm2_weights_fp4[i] permute_indices.to(gemm2_weights_fp4.device)
.view(torch.uint8)[permute_indices.to(gemm2_weights_fp4.device)] ]
.contiguous()
)
permute_sf_indices = get_w2_permute_indices_with_cache( permute_sf_indices = get_w2_permute_indices_with_cache(
_cache_permute_indices, _cache_permute_indices,
@@ -738,30 +754,24 @@ def prepare_static_weights_for_trtllm_fp4_moe(
epilogue_tile_m, epilogue_tile_m,
num_elts_per_sf=16, num_elts_per_sf=16,
) )
gemm2_scales_fp4_shuffled.append( torch.index_select(
nvfp4_block_scale_interleave( gemm2_scales_linear_fp4[i].view(torch.uint8),
gemm2_scales_linear_fp4[i] 0,
.view(torch.uint8)[ permute_sf_indices.to(gemm2_scales_linear_fp4.device),
permute_sf_indices.to(gemm2_scales_linear_fp4.device) out=g2s_scratch,
]
.contiguous()
)
) )
gemm2_scales_fp4_shuffled[i] = nvfp4_block_scale_interleave(g2s_scratch)
# Stack weights for all experts del g1s_scratch, g2s_scratch
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)
)
gemm2_weights_fp4_shuffled = torch.stack(gemm2_weights_fp4_shuffled) # Weight outputs stay as uint8 (FP4 packed) — the TRTLLM kernel expects this.
gemm2_scales_fp4_shuffled = ( gemm1_scales_fp4_shuffled = gemm1_scales_fp4_shuffled.view(
torch.stack(gemm2_scales_fp4_shuffled) torch.float8_e4m3fn
.view(torch.float8_e4m3fn) ).reshape(num_experts, gemm1_rows, hidden_size // 16)
.reshape(num_experts, hidden_size, intermediate_size // 16)
) gemm2_scales_fp4_shuffled = gemm2_scales_fp4_shuffled.view(
torch.float8_e4m3fn
).reshape(num_experts, hidden_size, intermediate_size // 16)
return ( return (
gemm1_weights_fp4_shuffled, gemm1_weights_fp4_shuffled,
gemm1_scales_fp4_shuffled, gemm1_scales_fp4_shuffled,