[loader] Reduce transient allocations in NVFP4 MoE setup (#26861)
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user