[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)
|
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
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user