Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
173 changes: 173 additions & 0 deletions tests/model_executor/model_loader/test_reload.py
Original file line number Diff line number Diff line change
Expand Up @@ -444,6 +444,179 @@ def process_weights_after_loading(self, layer):
assert torch.equal(layer.g_idx_sort_indices.data, expected_sort_indices)


_FLASHINFER_NUM_EXPERTS = 8

# variant -> (quant_dtype, weight_dtype, gemm1_clamp_limit,
# expected {constant name: value})
_FLASHINFER_CONSTANT_VARIANTS = {
"mxfp4_bf16": (
None,
"mxfp4",
None,
{"gemm1_alpha": 1.702, "gemm1_beta": 1.0, "gemm1_clamp_limit": 7.0},
),
"mxfp4_mxfp8": (
"mxfp8",
"mxfp4",
None,
{
"gemm1_alpha": 1.702,
"gemm1_beta": 1.0,
"gemm1_clamp_limit": 7.0,
"fake_input_scale": 1.0,
},
),
"fp8_clamp": (
torch.float8_e4m3fn,
torch.float8_e4m3fn,
3.0,
{"gemm1_clamp_limit": 3.0},
),
}


def _make_flashinfer_experts(layer, variant):
from vllm.model_executor.layers.fused_moe.activation import MoEActivation
from vllm.model_executor.layers.fused_moe.config import (
FusedMoEConfig,
FusedMoEParallelConfig,
FusedMoEQuantConfig,
RoutingMethodType,
)
from vllm.model_executor.layers.fused_moe.experts.flashinfer_cutlass_moe import (
FlashInferExperts,
)

quant_dtype, weight_dtype, gemm1_clamp_limit, _ = _FLASHINFER_CONSTANT_VARIANTS[
variant
]
moe_config = FusedMoEConfig(
num_experts=_FLASHINFER_NUM_EXPERTS,
experts_per_token=2,
hidden_dim=128,
intermediate_size=256,
num_local_experts=_FLASHINFER_NUM_EXPERTS,
num_logical_experts=_FLASHINFER_NUM_EXPERTS,
activation=MoEActivation.SILU,
device="cpu",
routing_method=RoutingMethodType.TopK,
moe_parallel_config=FusedMoEParallelConfig(
tp_size=1,
pcp_size=1,
dp_size=1,
ep_size=1,
tp_rank=0,
pcp_rank=0,
dp_rank=0,
ep_rank=0,
sp_size=1,
use_ep=False,
all2all_backend="naive",
enable_eplb=False,
),
in_dtype=torch.bfloat16,
)
quant_config = FusedMoEQuantConfig.make(
quant_dtype=quant_dtype,
weight_dtype=weight_dtype,
gemm1_clamp_limit=gemm1_clamp_limit,
)
return FlashInferExperts(moe_config, quant_config, layer=layer)


@pytest.mark.parametrize("variant", sorted(_FLASHINFER_CONSTANT_VARIANTS))
def test_flashinfer_cutlass_moe_constants_preserve_addresses(variant, dist_init):
"""FlashInferExperts is rebuilt by every post-load pass, but its per-expert
SwiGLU constants are read by captured CUDA graphs, so each rebuild must
reuse the storage the graph captured instead of allocating fresh tensors."""
expected = _FLASHINFER_CONSTANT_VARIANTS[variant][3]

layer = torch.nn.Module()
experts = _make_flashinfer_experts(layer, variant)

pointers = {name: getattr(layer, name).data_ptr() for name in expected}

# Reload: the quant method rebuilds the kernel, constructing a new
# experts object against the same layer
experts = _make_flashinfer_experts(layer, variant)

for name, value in expected.items():
constant = getattr(layer, name)
assert constant.data_ptr() == pointers[name], name
# registered as a Parameter so layerwise reload copy-back preserves it
assert isinstance(constant, torch.nn.Parameter), name
# the rebuilt experts object reads from the preserved storage
assert getattr(experts, name) is constant, name
assert torch.equal(
constant.data,
torch.full((_FLASHINFER_NUM_EXPERTS,), value, dtype=torch.float32),
), name


def test_flashinfer_cutlass_moe_layerwise_reload_accounting(dist_init):
"""The SwiGLU constants are generated during weight processing and never
loaded from checkpoints. Registering them as Parameters must not count
them toward `load_numel_total`: reload restores the construction-time
tensor set before sizing, so the layer still processes during streaming
instead of deferring (and buffering weights) until finalization."""
from vllm.model_executor.layers.quantization.base_config import (
QuantizeMethodBase,
)
from vllm.model_executor.model_loader.reload.layerwise import get_layerwise_info

class _KernelQuantMethod(QuantizeMethodBase):
def create_weights(self, layer, *args, **kwargs):
raise NotImplementedError

def apply(self, layer, *args, **kwargs):
raise NotImplementedError

def process_weights_after_loading(self, layer):
# The kernel (and its experts object) is rebuilt on every pass
self.experts = _make_flashinfer_experts(layer, "mxfp4_bf16")

def _load_checkpoint_format_weights(layer, checkpoint):
for name, weight in checkpoint.items():
param = torch.nn.Parameter(weight.clone(), requires_grad=False)
param.weight_loader = default_weight_loader
setattr(layer, name, param)

def _make_checkpoint(fill):
return {
"w13_weight": torch.full((8, 16, 4), fill, dtype=torch.bfloat16),
"w2_weight": torch.full((8, 4, 8), fill, dtype=torch.bfloat16),
}

layer = torch.nn.Module()
layer.quant_method = _KernelQuantMethod()
_load_checkpoint_format_weights(layer, _make_checkpoint(1.0))

# Metadata is recorded at model construction, before any processing
record_metadata_for_reloading(layer)
checkpoint_numel = sum(t.numel() for t in get_layer_tensors(layer).values())

layer.quant_method.process_weights_after_loading(layer)
constants = {
name: getattr(layer, name)
for name in ("gemm1_alpha", "gemm1_beta", "gemm1_clamp_limit")
}

initialize_layerwise_reload(layer)
info = get_layerwise_info(layer)
assert info.load_numel_total == checkpoint_numel

# Stream a new checkpoint; the layer must process as soon as its last
# tensor arrives
for name, weight in _make_checkpoint(2.0).items():
param = getattr(layer, name)
param.weight_loader(param, weight)

assert not info.can_load()
assert not info.loaded_weights
for name, constant in constants.items():
assert getattr(layer, name) is constant, name


def test_model_cleanup(dist_init, default_vllm_config):
layer = QKVParallelLinear(2, 3, 4)
assert layer.weight.weight_loader.__self__ is layer
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
kNvfp4Dynamic,
kNvfp4Static,
)
from vllm.model_executor.utils import replace_parameter
from vllm.platforms import current_platform
from vllm.utils.flashinfer import (
flashinfer_cutlass_fused_moe,
Expand Down Expand Up @@ -69,6 +70,7 @@ def __init__(
self,
moe_config: mk.FusedMoEConfig,
quant_config: FusedMoEQuantConfig,
layer: torch.nn.Module | None = None,
):
super().__init__(moe_config, quant_config)

Expand All @@ -95,34 +97,46 @@ def __init__(
self.use_deepseek_fp8_block_scale = quant_config.is_block_quantized
self.gemm1_clamp_limit: torch.Tensor | None = None
if quant_config.gemm1_clamp_limit is not None:
self.gemm1_clamp_limit = torch.tensor(
[quant_config.gemm1_clamp_limit] * self.num_experts,
dtype=torch.float32,
device=self.device,
self.gemm1_clamp_limit = self._stable_constant(
layer, "gemm1_clamp_limit", quant_config.gemm1_clamp_limit
)

if quant_config.weight_quant_dtype == "mxfp4":
# This value is used specifically for gpt-oss,
# Need to revisit this for other models
self.gemm1_alpha = torch.tensor(
[1.702] * self.num_experts, dtype=torch.float32, device=self.device
)
self.gemm1_beta = torch.tensor(
[1.0] * self.num_experts, dtype=torch.float32, device=self.device
)
self.gemm1_alpha = self._stable_constant(layer, "gemm1_alpha", 1.702)
self.gemm1_beta = self._stable_constant(layer, "gemm1_beta", 1.0)
if self.gemm1_clamp_limit is None:
self.gemm1_clamp_limit = torch.tensor(
[7.0] * self.num_experts,
dtype=torch.float32,
device=self.device,
self.gemm1_clamp_limit = self._stable_constant(
layer, "gemm1_clamp_limit", 7.0
)
if quant_config.quant_dtype == "mxfp8":
self.fake_input_scale = torch.ones(
self.num_experts,
device=self.device,
dtype=torch.float32,
self.fake_input_scale = self._stable_constant(
layer, "fake_input_scale", 1.0
)

def _stable_constant(
self, layer: torch.nn.Module | None, name: str, value: float
) -> torch.Tensor:
"""Allocate a per-expert constant at a reload-stable address.

This object is rebuilt by every process_weights_after_loading pass,
but these tensors are passed to the kernel each forward, so captured
CUDA graphs bake their addresses (see #48312). Registering them on
the layer with replace_parameter(prefer_copy=True) makes every
rebuild copy into the storage the graph captured instead of
allocating a fresh tensor.
"""
constant = torch.full(
(self.num_experts,), value, dtype=torch.float32, device=self.device
)
if layer is None:
# Direct construction without an owning layer (tests, benchmarks):
# nothing survives a rebuild, so plain storage is equivalent.
return constant
replace_parameter(layer, name, constant, prefer_copy=True)
return getattr(layer, name)

@property
def expects_unquantized_inputs(self) -> bool:
return self.quant_config.use_fp8_w8a8 and self.quant_config.is_block_quantized
Expand Down
9 changes: 9 additions & 0 deletions vllm/model_executor/layers/fused_moe/oracle/fp8.py
Original file line number Diff line number Diff line change
Expand Up @@ -687,10 +687,19 @@ def make_fp8_moe_kernel(

logger.info_once("Using %s", prepare_finalize.__class__.__name__)

from vllm.model_executor.layers.fused_moe.experts.flashinfer_cutlass_moe import (
FlashInferExperts,
)

extra_kwargs = {}
if fp8_backend == Fp8MoeBackend.HUMMING:
assert layer is not None
extra_kwargs = {"layer": layer}
elif experts_cls is FlashInferExperts:
# FlashInferExperts registers its per-expert constants on the layer
# so they keep their storage across kernel rebuilds (see #48312).
assert layer is not None
extra_kwargs = {"layer": layer}

# Create Experts.
if prepare_finalize.activation_format == mk.FusedMoEActivationFormat.BatchedExperts:
Expand Down
9 changes: 9 additions & 0 deletions vllm/model_executor/layers/fused_moe/oracle/mxfp4.py
Original file line number Diff line number Diff line change
Expand Up @@ -1719,10 +1719,19 @@ def make_mxfp4_moe_kernel(
logger.info_once("Using %s", prepare_finalize.__class__.__name__)
logger.info_once("Using %s", experts_cls.__name__)

from vllm.model_executor.layers.fused_moe.experts.flashinfer_cutlass_moe import (
FlashInferExperts,
)

extra_kwargs = {}
if mxfp4_backend == Mxfp4MoeBackend.HUMMING:
assert layer is not None
extra_kwargs["layer"] = layer
elif experts_cls is FlashInferExperts:
# FlashInferExperts registers its per-expert constants on the layer
# so they keep their storage across kernel rebuilds (see #48312).
assert layer is not None
extra_kwargs["layer"] = layer

# Create Experts.
if prepare_finalize.activation_format == mk.FusedMoEActivationFormat.BatchedExperts:
Expand Down
9 changes: 9 additions & 0 deletions vllm/model_executor/layers/fused_moe/oracle/nvfp4.py
Original file line number Diff line number Diff line change
Expand Up @@ -543,10 +543,19 @@ def make_nvfp4_moe_kernel(

logger.info_once("Using %s", prepare_finalize.__class__.__name__)

from vllm.model_executor.layers.fused_moe.experts.flashinfer_cutlass_moe import (
FlashInferExperts,
)

extra_kwargs = {}
if backend == NvFp4MoeBackend.HUMMING:
assert layer is not None
extra_kwargs = {"layer": layer}
elif experts_cls is FlashInferExperts:
# FlashInferExperts registers its per-expert constants on the layer
# so they keep their storage across kernel rebuilds (see #48312).
assert layer is not None
extra_kwargs = {"layer": layer}
if backend == NvFp4MoeBackend.FLASHINFER_TRTLLM and per_token_activation:
extra_kwargs["per_token_activation"] = True

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -215,6 +215,7 @@ def process_weights_after_loading(self, layer: RoutedExperts) -> None:
experts_cls=self.experts_cls,
mxfp4_backend=self.mxfp4_backend,
routing_tables=layer._expert_routing_tables(),
layer=layer,
)

def apply(
Expand Down
2 changes: 2 additions & 0 deletions vllm/model_executor/layers/quantization/quark/quark_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -457,6 +457,7 @@ def _setup_kernel(self, layer: RoutedExperts) -> None:
fp8_backend=self.fp8_backend,
experts_cls=self.experts_cls,
routing_tables=layer._expert_routing_tables(),
layer=layer,
)

def get_fused_moe_quant_config(self, layer: RoutedExperts) -> FusedMoEQuantConfig:
Expand Down Expand Up @@ -1339,6 +1340,7 @@ def _setup_kernel(self, layer: RoutedExperts):
mxfp4_backend=self.mxfp4_backend,
experts_cls=self.experts_cls,
routing_tables=layer._expert_routing_tables(),
layer=layer,
)

def get_fused_moe_quant_config(
Expand Down
Loading