fix(nvfp4): make process_weights_after_loading hot-reload-safe via alias-when-same-shape (#25190)

Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-05-13 16:57:39 -07:00
committed by GitHub
co-authored by Claude Opus 4.7
parent 37f18438c5
commit 6c0633b0b1
2 changed files with 59 additions and 19 deletions
@@ -57,7 +57,7 @@ from sglang.srt.layers.quantization.utils import (
swizzle_blockscale,
)
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.utils import copy_or_rebind_param
from sglang.srt.layers.utils import alias_or_bind_derived_param, copy_or_rebind_param
from sglang.srt.utils.common import (
is_cuda,
is_sm120_supported,
@@ -1363,13 +1363,15 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
input_scale_2 = layer.input_scale.max().to(torch.float32)
weight_scale_2 = layer.weight_scale_2.max().to(torch.float32)
# alpha / input_scale_inv stay as scalar Parameters. Aliasing them into
# the [N_partitions] source slot breaks fused-QKV linears whose
# downstream kernels assume scalar input scale.
copy_or_rebind_param(
layer, "alpha", (input_scale_2 * weight_scale_2).to(torch.float32)
)
copy_or_rebind_param(
layer, "input_scale_inv", (1 / input_scale_2).to(torch.float32)
)
del layer.input_scale, layer.weight_scale_2
# Store original output size before any padding
layer.output_size_per_partition = layer.weight.shape[0]
@@ -1420,10 +1422,11 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
.view(torch.float8_e4m3fn)
)
copy_or_rebind_param(layer, "weight_scale_interleaved", scale)
alias_or_bind_derived_param(
layer, "weight_scale", "weight_scale_interleaved", scale
)
copy_or_rebind_param(layer, "weight", weight)
layer.weights_padding_cols = weights_padding_cols
del layer.weight_scale
return
# Pad weights for CUTLASS/FlashInfer kernel alignment (K and N divisible by 32)
@@ -1453,8 +1456,9 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
if scale_ndim == 2
else padded_scales.reshape(B, M_padded, K_padded)
)
copy_or_rebind_param(layer, "weight_scale_interleaved", padded_scales)
del layer.weight_scale
alias_or_bind_derived_param(
layer, "weight_scale", "weight_scale_interleaved", padded_scales
)
def apply(
self,
@@ -1704,7 +1708,6 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
Only supports pre-quantized checkpoints with FP8 weights and scales.
"""
# GEMM 1 scale processing
if layer.moe_runner_config.is_gated:
if layer.w13_weight_scale_2.dim() == 1:
@@ -1777,12 +1780,6 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
"w2_input_scale_quant",
(1 / w2_input_scale).to(torch.float32),
)
del layer.w13_input_scale, layer.w2_input_scale
# TODO: w13_weight_scale_2 / w2_weight_scale_2 are also unused by apply()
# after this point. flashinfer_cutedsl reads them via hasattr() but has a
# mathematically-equivalent fallback through w13_input_scale_quant and
# g1_alphas. Kept for now to avoid a silent code-path switch and possible
# sub-ULP precision drift; revisit once the fallback is validated.
# TODO: for flashinfer always do MOE_NVFP4_DISPATCH
layer.dispatcher.set_quant_config(
@@ -1835,7 +1832,10 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
# FlashInfer TRTLLM processing - handles both w13 and w2
align_fp4_moe_weights_for_flashinfer_trtllm(layer)
del layer.w13_blockscale_swizzled, layer.w2_blockscale_swizzled
# TRTLLM doesn't read *_blockscale_swizzled; alias to free the
# placeholders from create_weights.
layer.w13_blockscale_swizzled = layer.w13_weight_scale
layer.w2_blockscale_swizzled = layer.w2_weight_scale
else:
# CUTLASS processing - handle w13 and w2 separately
@@ -1864,8 +1864,11 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
# Process w13 weights
w13_blockscale_swizzled = swizzle_blockscale(layer.w13_weight_scale)
copy_or_rebind_param(
layer, "w13_blockscale_swizzled", w13_blockscale_swizzled
alias_or_bind_derived_param(
layer,
"w13_weight_scale",
"w13_blockscale_swizzled",
w13_blockscale_swizzled,
)
w13_weight = layer.w13_weight
@@ -1902,8 +1905,11 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
# Process w2 weights
w2_blockscale_swizzled = swizzle_blockscale(layer.w2_weight_scale)
copy_or_rebind_param(
layer, "w2_blockscale_swizzled", w2_blockscale_swizzled
alias_or_bind_derived_param(
layer,
"w2_weight_scale",
"w2_blockscale_swizzled",
w2_blockscale_swizzled,
)
if self._is_cutedsl_v2_standard:
@@ -1969,7 +1975,6 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
intermediate_size_per_partition=inter_size, # n
hidden_size=hidden_size,
) # k
del layer.w13_weight_scale, layer.w2_weight_scale
@property
def load_up_proj_weight_first(self) -> bool:
+35
View File
@@ -71,6 +71,41 @@ def copy_or_rebind_param(
setattr(module, name, Parameter(new_value, requires_grad=False))
def alias_or_bind_derived_param(
module: torch.nn.Module,
source_name: str,
derived_name: str,
derived_value: torch.Tensor,
) -> None:
"""Bind a post-processed (derived) tensor to a derived attribute name.
When `derived_value` is broadcastable to the source Parameter's shape (and
dtype matches), write it broadcast-filled into the source's storage in
place and register `derived_name` as an alias of the source Parameter. The
two attribute names then share one underlying buffer, so:
- apply() can read via `derived_name`
- update_weights_from_disk can keep refilling `source_name` (the loader
re-runs process_weights_after_loading which re-derives in place)
- peak GPU memory is the source size, not source + derived.
When the shapes are not broadcast-compatible, fall back to allocating a
separate Parameter under `derived_name` via copy_or_rebind_param.
"""
derived_value = derived_value.detach()
source = getattr(module, source_name, None)
if isinstance(source, Parameter) and source.data.dtype == derived_value.dtype:
try:
broadcast = torch.broadcast_to(derived_value, source.data.shape)
except RuntimeError:
broadcast = None
if broadcast is not None:
source.data.copy_(broadcast)
source.requires_grad_(False)
setattr(module, derived_name, source)
return
copy_or_rebind_param(module, derived_name, derived_value)
class PPMissingLayer(torch.nn.Identity):
# Adapted from
# https://github.com/vllm-project/vllm/blob/18ed3132d2bfe1df9a74729457b69243955221e8/vllm/model_executor/models/utils.py#L468C1-L486C1