From b88aa2c71c9e9b06cc4c60054ef50ce7ad89e13a Mon Sep 17 00:00:00 2001 From: sna Date: Tue, 21 Jul 2026 16:16:02 -0700 Subject: [PATCH 01/76] perf(vllm): optimize quantized refit paths Reduce MXFP8 and ModelOpt refit overhead while preserving transport and checkpoint-engine lifecycle correctness. Signed-off-by: sna --- examples/configs/grpo_math_1B.yaml | 10 + nemo_rl/algorithms/distillation.py | 6 +- nemo_rl/algorithms/grpo.py | 5 +- nemo_rl/algorithms/ppo.py | 10 +- nemo_rl/algorithms/utils.py | 37 +- .../models/generation/vllm_quant_backend.py | 14 +- nemo_rl/models/generation/interfaces.py | 12 +- .../generation/vllm/checkpoint_engine.py | 24 +- nemo_rl/models/generation/vllm/config.py | 5 + .../generation/vllm/quantization/fp8.py | 550 ++++++++++++++++-- .../vllm/quantization/fp8_train_utils.py | 82 +++ .../models/generation/vllm/vllm_backend.py | 195 ++++++- .../models/generation/vllm/vllm_generation.py | 21 +- nemo_rl/models/generation/vllm/vllm_worker.py | 16 +- .../generation/vllm/vllm_worker_async.py | 12 +- nemo_rl/models/megatron/setup.py | 20 +- nemo_rl/models/policy/__init__.py | 16 + nemo_rl/models/policy/interfaces.py | 12 + nemo_rl/models/policy/lm_policy.py | 14 + nemo_rl/models/policy/utils.py | 43 +- .../policy/workers/dtensor_policy_worker.py | 8 + .../workers/dtensor_policy_worker_v2.py | 8 + .../policy/workers/megatron_policy_worker.py | 227 ++++++-- pyrefly.toml | 1 + .../models/generation/test_mxfp8_prequant.py | 149 +++++ .../models/generation/test_vllm_backend.py | 2 - .../generation/test_vllm_checkpoint_engine.py | 21 +- .../test_vllm_modelopt_real_quant_config.py | 2 +- .../unit/reference_configs/grpo_math_1B.yaml | 3 + 29 files changed, 1340 insertions(+), 185 deletions(-) create mode 100644 tests/unit/models/generation/test_mxfp8_prequant.py diff --git a/examples/configs/grpo_math_1B.yaml b/examples/configs/grpo_math_1B.yaml index 88a946c61b0..1cb3110fd4c 100644 --- a/examples/configs/grpo_math_1B.yaml +++ b/examples/configs/grpo_math_1B.yaml @@ -157,6 +157,12 @@ policy: megatron_cfg: enabled: false + # When true, offload_after_refit skips rerunning the full offload_before_refit + # pass and only re-offloads the optimizer with a single allocator cleanup. + refit_slim_offload_after: false + # When true, stage the per-step reference-policy swap in persistent pinned + # CPU buffers instead of pageable copies (faster PCIe, more host RAM). + pinned_reference_swap: false checkpoint: async_save: true ckpt_assume_constant_structure: true @@ -305,6 +311,10 @@ policy: # makes the training sequence length divisible by the tensor parallel size # this is useful for sequence parallel training make_sequence_length_divisible_by: ${policy.dtensor_cfg.tensor_parallel_size} + # Keep refit CUDA-IPC staging buffers allocated across refits (skips two large + # allocations plus a gc/empty_cache pair per refit). Best with a fixed + # refit_buffer_size_gb; buffers stay resident on the trainer GPU between refits. + refit_persistent_ipc_buffers: false max_grad_norm: 1.0 optimizer: diff --git a/nemo_rl/algorithms/distillation.py b/nemo_rl/algorithms/distillation.py index 6f536ed616e..7f0442860a9 100644 --- a/nemo_rl/algorithms/distillation.py +++ b/nemo_rl/algorithms/distillation.py @@ -37,7 +37,7 @@ DistillationLossDataDict, DistillationLossFn, ) -from nemo_rl.algorithms.utils import set_seed +from nemo_rl.algorithms.utils import maybe_enable_refit_prequantize, set_seed from nemo_rl.data import DataConfig from nemo_rl.data.collate_fn import rl_collate_fn from nemo_rl.data.datasets import AllTaskProcessedDataset @@ -619,7 +619,9 @@ def init_nemo_gym(): student_generation.weight_synchronizer.init_communicator() elif student_generation is not None: state_dict_info = student_policy.prepare_refit_info() - student_generation.prepare_refit_info(state_dict_info) + maybe_enable_refit_prequantize( + student_policy, student_generation, state_dict_info, master_config.policy + ) # if it is not colocated inference, initialize collective communication for update weights if not colocated_inference and checkpoint_engine_config is None: diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index f5cb34adc19..96fee38314e 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -55,6 +55,7 @@ calculate_baseline_and_std_per_prompt, get_gdpo_reward_component_keys, log_generation_metrics_to_wandb, + maybe_enable_refit_prequantize, print_efficiency_summary, print_performance_metrics, set_seed, @@ -1367,7 +1368,9 @@ def init_vllm_then_policy(): else: state_dict_info = policy.prepare_refit_info() if policy_generation is not None: - policy_generation.prepare_refit_info(state_dict_info) + maybe_enable_refit_prequantize( + policy, policy_generation, state_dict_info, master_config.policy + ) # Spin up non-colocated OPD teacher worker groups AFTER policy / vLLM are # ready. Parallelizing with policy init races on Megatron-Bridge's HF->mcore diff --git a/nemo_rl/algorithms/ppo.py b/nemo_rl/algorithms/ppo.py index 9971fba15ef..62f6a546074 100644 --- a/nemo_rl/algorithms/ppo.py +++ b/nemo_rl/algorithms/ppo.py @@ -48,7 +48,11 @@ RewardShapingConfig, apply_reward_shaping, ) -from nemo_rl.algorithms.utils import print_performance_metrics, set_seed +from nemo_rl.algorithms.utils import ( + maybe_enable_refit_prequantize, + print_performance_metrics, + set_seed, +) from nemo_rl.data import DataConfig from nemo_rl.data.collate_fn import rl_collate_fn from nemo_rl.data.datasets import AllTaskProcessedDataset @@ -654,7 +658,9 @@ def initialize_generation_with_policy( # prepare refit info state_dict_info = policy.prepare_refit_info() if policy_generation is not None: - policy_generation.prepare_refit_info(state_dict_info) + maybe_enable_refit_prequantize( + policy, policy_generation, state_dict_info, master_config.policy + ) # Calculate total setup time total_setup_time = time.perf_counter() - setup_start_time diff --git a/nemo_rl/algorithms/utils.py b/nemo_rl/algorithms/utils.py index 41794b243d5..0f83732eea4 100644 --- a/nemo_rl/algorithms/utils.py +++ b/nemo_rl/algorithms/utils.py @@ -16,7 +16,7 @@ import random import warnings from functools import partial, wraps -from typing import Any, Optional +from typing import TYPE_CHECKING, Any, Optional import numpy as np import torch @@ -27,9 +27,15 @@ ) from nemo_rl.data.chat_templates import COMMON_CHAT_TEMPLATES -from nemo_rl.models.policy import TokenizerConfig +from nemo_rl.models.policy import PolicyConfig, TokenizerConfig from nemo_rl.utils.logger import Logger +if TYPE_CHECKING: + # Runtime import would cycle: policy.interfaces pulls in algorithms.loss, + # which imports back into algorithms.utils. + from nemo_rl.models.generation.interfaces import GenerationInterface + from nemo_rl.models.policy.interfaces import ColocatablePolicyInterface + def get_gdpo_reward_component_keys(batch) -> list[str]: """Return batch keys that are named reward components (e.g. reward/correctness) in sorted order.""" @@ -1029,3 +1035,30 @@ def print_efficiency_summary( loggable["efficiency/total_wall_time_s"] = total_wall_time_s return loggable + + +def maybe_enable_refit_prequantize( + policy: "ColocatablePolicyInterface", + policy_generation: "GenerationInterface", + state_dict_info: Optional[dict[str, Any]], + policy_config: PolicyConfig, +) -> None: + """Complete the trainer-side pre-quantized refit handshake if requested. + + The generation backend's prepare_refit_info returns the fp8-eligible + parameter names when vllm_cfg.refit_prequantize is enabled; the trainer + then quantizes exactly those during refit and the receiver's unpack + metadata is refreshed to the quantized dtypes plus scale entries. + """ + prequant_names = policy_generation.prepare_refit_info(state_dict_info) + if not prequant_names: + return + megatron_cfg = policy_config.get("megatron_cfg") + if not (megatron_cfg and megatron_cfg["enabled"]): + raise ValueError( + "vllm_cfg.refit_prequantize requires the Megatron policy backend " + "(policy.megatron_cfg.enabled=true); the DTensor workers do not " + "implement trainer-side pre-quantized refit." + ) + updated_info = policy.enable_refit_prequantize(prequant_names) + policy_generation.prepare_refit_info(updated_info) diff --git a/nemo_rl/modelopt/models/generation/vllm_quant_backend.py b/nemo_rl/modelopt/models/generation/vllm_quant_backend.py index 4f436492321..d496b981270 100644 --- a/nemo_rl/modelopt/models/generation/vllm_quant_backend.py +++ b/nemo_rl/modelopt/models/generation/vllm_quant_backend.py @@ -16,7 +16,7 @@ import types from collections.abc import Iterator from contextlib import ExitStack, contextmanager -from typing import Any +from typing import Any, Optional import torch import vllm # noqa: F401 @@ -510,10 +510,15 @@ def _synchronize_before_ipc_data_ack(self) -> None: return super()._synchronize_before_ipc_data_ack() - def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: - super().prepare_refit_info(state_dict_info) + def prepare_refit_info( + self, state_dict_info: dict[str, Any] + ) -> Optional[list[str]]: if not self._is_real_quant_model(): - return + return super().prepare_refit_info(state_dict_info) + + # Real quantization owns a separate refit handshake and must not import + # the legacy FP8 quantization path. + self.state_dict_info = state_dict_info self._get_modelopt_reload_roots() quant_config = ( self.model_runner.vllm_config.model_config.hf_config.quantization_config @@ -532,6 +537,7 @@ def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: "Fused ModelOpt MoE refits require all experts local; " "vLLM expert parallelism is unsupported" ) + return None @contextmanager def _patch_named_parameters_to_include_buffers(self, model): diff --git a/nemo_rl/models/generation/interfaces.py b/nemo_rl/models/generation/interfaces.py index 6a76cb40861..19850611938 100644 --- a/nemo_rl/models/generation/interfaces.py +++ b/nemo_rl/models/generation/interfaces.py @@ -330,8 +330,16 @@ def requires_kv_scale_sync(self) -> bool: """Whether the generation backend requires KV cache scales synchronization.""" return False - def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: - """Prepare the info for refit.""" + def prepare_refit_info( + self, state_dict_info: Optional[dict[str, Any]] + ) -> Optional[list[str]]: + """Prepare the info for refit. + + Returns: + Optionally, the parameter names the backend wants pre-quantized on + the trainer before streaming (e.g. vllm_cfg.refit_prequantize); + None when no trainer-side pre-quantization is requested. + """ raise NotImplementedError def update_weights_via_ipc_zmq(self) -> list[ray.ObjectRef]: diff --git a/nemo_rl/models/generation/vllm/checkpoint_engine.py b/nemo_rl/models/generation/vllm/checkpoint_engine.py index 4a9a42bdba5..e1e49ad994f 100644 --- a/nemo_rl/models/generation/vllm/checkpoint_engine.py +++ b/nemo_rl/models/generation/vllm/checkpoint_engine.py @@ -143,18 +143,18 @@ async def _update_weights_from_checkpoint_engine_async(self) -> bool: load_time = 0.0 start_time = time.time() - async for weight_batch in self.checkpoint_engine.receive_weight_batches(): - loaded_batches += 1 - loaded_tensors += len(weight_batch) - loaded_bytes += sum(weight.nbytes for _name, weight in weight_batch) - - load_start = time.time() - self._load_weights(weight_batch) - torch.cuda.current_stream().synchronize() - load_time += time.time() - load_start - del weight_batch - - self._maybe_process_fp8_kv_cache() + with self._weight_update_lifecycle("checkpoint_engine") as finalize: + async for weight_batch in self.checkpoint_engine.receive_weight_batches(): + loaded_batches += 1 + loaded_tensors += len(weight_batch) + loaded_bytes += sum(weight.nbytes for _name, weight in weight_batch) + + load_start = time.time() + self._load_weights(weight_batch) + torch.cuda.current_stream().synchronize() + load_time += time.time() - load_start + del weight_batch + finalize() total_time = time.time() - start_time loaded_gib = loaded_bytes / (1024 * 1024 * 1024) diff --git a/nemo_rl/models/generation/vllm/config.py b/nemo_rl/models/generation/vllm/config.py index 03353836ca6..9d3295b3460 100644 --- a/nemo_rl/models/generation/vllm/config.py +++ b/nemo_rl/models/generation/vllm/config.py @@ -37,6 +37,11 @@ class VllmSpecificArgs(TypedDict): precision: NotRequired[str] # Use ModelOpt MXFP8 quantization when precision is fp8. is_mx: NotRequired[bool] + # With is_mx, quantize weights to MXFP8 on the trainer during refit and + # stream E4M3 data plus scales (~47% smaller payload) instead of BF16; + # the vLLM worker then skips its per-refit re-quantization. Requires the + # Megatron policy backend. + refit_prequantize: NotRequired[bool] kv_cache_dtype: Literal["auto", "fp8", "fp8_e4m3"] enforce_eager: NotRequired[bool] enable_return_routed_experts: NotRequired[bool] diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index 4f0e3b1f70a..8dd9ede0ccf 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -53,6 +53,9 @@ class FP8Config: kv_cache_dtype: str = "auto" use_fp8_weights: bool = True # Whether model weights are quantized to FP8 is_mx: bool = False + # Weights arrive from the trainer already MXFP8-quantized (E4M3 data plus + # *_scale_from_checkpoint entries), so load_weights skips re-quantization. + refit_prequantize: bool = False @dataclass() @@ -159,6 +162,12 @@ def apply_fp8_patches(self, fp8_config): process_weights_after_loading_mxfp8_moe, ) ) + fp8_state.vllm_patches.append( + patch( + "vllm.model_executor.layers.quantization.modelopt.ModelOptMxFp8FusedMoE.apply_monolithic", + apply_monolithic_mxfp8_moe, + ) + ) # These patches add support for pow2, e8 dynamic activation scalings factors which are believed to have higher # SNR compared to plain fp32 scaling factors. This feature is still under active research. @@ -215,6 +224,7 @@ def init_fp8(vllm_cfg, model_name, model_parallel_size): } if is_mx: fp8_config_kwargs["is_mx"] = True + fp8_config_kwargs["refit_prequantize"] = bool(vllm_cfg.get("refit_prequantize")) if vllm_cfg.get("pow2_weight_scaling_factors") is False: raise ValueError("only pow2 weight scaling factors are supported for MXFP8") if vllm_cfg.get("pow2_activation_scaling_factors") is False: @@ -436,6 +446,11 @@ def load_weights(weights, model_runner): if not _is_fp8_weight(k, model): weights_quantized.append((k, v)) continue + if v.dtype == torch.float8_e4m3fn: + # Already quantized on the trainer (vllm_cfg.refit_prequantize); the + # matching *_scale_from_checkpoint entry arrives as its own weight. + weights_quantized.append([k, v]) + continue # Cast the weight into fp8 and its scale factor if global_fp8_config.is_mx: from vllm.model_executor.layers.quantization.utils.mxfp8_utils import ( @@ -450,13 +465,21 @@ def load_weights(weights, model_runner): ) param_scale = torch.squeeze(param_scale, dim=-1) if global_fp8_config.is_mx: + # All-zero blocks quantize to E8M0 byte 0, which destabilizes the + # TRTLLM MXFP8 kernel; clamp to byte 1 (weights are 0 anyway). + param_scale = torch.where( + param_scale == 0, torch.ones_like(param_scale), param_scale + ) weights_quantized.append([k, param_lp]) weights_quantized.append([k + "_scale_from_checkpoint", param_scale]) else: weights_quantized.append([k, param_lp]) weights_quantized.append([k + "_scale_inv", param_scale]) - # Finally load the weights into vllm - model.load_weights(weights_quantized) + # Finally load the weights into vllm. Deferred: importing vllm_backend at + # module top would cycle through the nemo_rl generation package init. + from nemo_rl.models.generation.vllm.vllm_backend import load_weights_maybe_cached + + load_weights_maybe_cached(model, weights_quantized) def cast_tensor_to_fp8_blockwise( @@ -841,56 +864,219 @@ def process_weights_after_loading_moe(self, layer) -> None: ) -def process_weights_after_loading_mxfp8_moe(self, layer) -> None: - """Shuffle weights and scales into FlashInfer TRTLLM MXFP8 layout.""" +def _round_up(value: int, multiple: int) -> int: + return ((value + multiple - 1) // multiple) * multiple + + +def _pad_tensor_dim( + tensor: torch.Tensor, dim: int, padded_size: int, pad_value: int = 0 +) -> torch.Tensor: + current_size = tensor.shape[dim] + if current_size == padded_size: + return tensor + padded_shape = list(tensor.shape) + padded_shape[dim] = padded_size + padded = torch.zeros(padded_shape, dtype=tensor.dtype, device=tensor.device) + if pad_value != 0: + padded.fill_(pad_value) + padded.narrow(dim, 0, current_size).copy_(tensor) + return padded + + +def _pad_w13_shards( + tensor: torch.Tensor, + intermediate_size_factor: int, + padded_intermediate_size: int, + pad_value: int = 0, +) -> torch.Tensor: + """Pad the intermediate dim of a [E, factor*I, K] tensor to [E, factor*I_pad, K]. + + Pads each of the factor shards separately so gated W13 halves stay aligned + for swap_w13_to_w31 and reorder_rows_for_gated_act_gemm. + """ + num_experts, total_rows, cols = tensor.shape + rows_per_shard = total_rows // intermediate_size_factor + if rows_per_shard == padded_intermediate_size: + return tensor + shards = tensor.reshape(num_experts, intermediate_size_factor, rows_per_shard, cols) + shards = _pad_tensor_dim(shards, 2, padded_intermediate_size, pad_value) + return shards.reshape( + num_experts, intermediate_size_factor * padded_intermediate_size, cols + ) + + +def _clamp_mxfp8_scale(scale: torch.Tensor) -> torch.Tensor: + return torch.where(scale == 0, torch.ones_like(scale), scale) + + +def _set_mxfp8_apply_tensor(layer, name: str, value: torch.Tensor) -> None: + existing = getattr(layer, name, None) + if existing is not None and existing.shape == value.shape: + # Keep storage stable across refits (CUDA graphs capture pointers). + existing.copy_(value) + else: + # Clone: value may view a scratch buffer shared across layers. + setattr(layer, name, torch.nn.Parameter(value.clone(), requires_grad=False)) + + +# Shared gather destinations keyed by (tag, shape, device). They persist across +# refits so the batched shuffle allocates nothing after the first pass; their +# contents are rewritten on every call, so a sleep-mode discard is harmless. +mxfp8_shuffle_scratch_buffers: dict[ + tuple[str, tuple[int, ...], torch.device], torch.Tensor +] = {} + +# One-shot flag for NRL_MXFP8_SHUFFLE_VERIFY: compare the batched shuffle +# against the per-expert reference on the first processed layer only. +mxfp8_shuffle_verified = False + + +def _mxfp8_scratch(tag: str, shape: torch.Size, device: torch.device) -> torch.Tensor: + key = (tag, tuple(shape), device) + buf = mxfp8_shuffle_scratch_buffers.get(key) + if buf is None: + buf = torch.empty(shape, dtype=torch.uint8, device=device) + mxfp8_shuffle_scratch_buffers[key] = buf + return buf + + +def _mxfp8_moe_row_permutations( + layer, + w13_weight: torch.Tensor, + w2_weight: torch.Tensor, + is_gated: bool, + epilogue_tile_m: int, +) -> tuple[torch.Tensor, torch.Tensor]: + """Composed row-index permutations for the batched TRTLLM MoE shuffle. + + shuffle_matrix_a / shuffle_matrix_sf_a and reorder_rows_for_gated_act_gemm + are input-independent row permutations (and the sf row indices equal the + weight row indices for the same row count), so composing the indices once + reproduces the per-expert call sequence as a single gather per tensor. + Cached on the layer as CPU tensors; device copies are made per call because + tensors allocated while vLLM loads the model land in the sleep-mode weights + pool, whose contents are discarded at sleep_level=2. + """ + perm_w13 = getattr(layer, "_mxfp8_shuffle_perm_w13", None) + perm_w2 = getattr(layer, "_mxfp8_shuffle_perm_w2", None) + if perm_w13 is None or perm_w2 is None: + from flashinfer.fused_moe.core import ( + get_reorder_rows_for_gated_act_gemm_row_indices, + ) + from flashinfer.utils import get_shuffle_matrix_a_row_indices + + perm_w13 = get_shuffle_matrix_a_row_indices(w13_weight[0], epilogue_tile_m) + if is_gated: + reorder = get_reorder_rows_for_gated_act_gemm_row_indices(w13_weight[0]) + perm_w13 = reorder[perm_w13] + perm_w2 = get_shuffle_matrix_a_row_indices(w2_weight[0], epilogue_tile_m) + layer._mxfp8_shuffle_perm_w13 = perm_w13 + layer._mxfp8_shuffle_perm_w2 = perm_w2 + device = w13_weight.device + return perm_w13.to(device), perm_w2.to(device) + + +def _shuffle_mxfp8_moe_batched( + layer, + w13_weight: torch.Tensor, + w2_weight: torch.Tensor, + w13_scale: torch.Tensor, + w2_scale: torch.Tensor, + is_gated: bool, + epilogue_tile_m: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Apply the TRTLLM row shuffles as one batched gather per stacked tensor. + + Bit-identical to the per-expert loop (`_shuffle_mxfp8_moe_per_expert`): + the row gathers run as a single dim-1 index_select over the whole + [E, M, K] tensor and the scale swizzle as one 3D block_scale_interleave, + whose flat output equals the stacked per-expert 2D outputs. Gathers write + into shared scratch buffers; callers copy_ into the persistent + destinations (w13/w2 weights may alias their own gather source). + """ + from flashinfer import block_scale_interleave + from vllm.model_executor.layers.quantization.utils.mxfp8_utils import ( + MXFP8_SCALE_DTYPE, + MXFP8_VALUE_DTYPE, + ) + + perm_w13, perm_w2 = _mxfp8_moe_row_permutations( + layer, w13_weight, w2_weight, is_gated, epilogue_tile_m + ) + num_experts = w13_weight.shape[0] + + w13_u8 = w13_weight.view(torch.uint8) + w2_u8 = w2_weight.view(torch.uint8) + w13_shuffled = torch.index_select( + w13_u8, 1, perm_w13, out=_mxfp8_scratch("w13", w13_u8.shape, w13_u8.device) + ) + w2_shuffled = torch.index_select( + w2_u8, 1, perm_w2, out=_mxfp8_scratch("w2", w2_u8.shape, w2_u8.device) + ) + + w13_sf_u8 = pad_flashinfer_scale_k(w13_scale.view(torch.uint8)) + w2_sf_u8 = pad_flashinfer_scale_k(w2_scale.view(torch.uint8)) + # Same constraint shuffle_matrix_sf_a asserts on the per-expert path. + assert w13_sf_u8.shape[1] % 128 == 0 and w2_sf_u8.shape[1] % 128 == 0 + w13_sf_gathered = torch.index_select( + w13_sf_u8, + 1, + perm_w13, + out=_mxfp8_scratch("w13_sf", w13_sf_u8.shape, w13_sf_u8.device), + ) + w2_sf_gathered = torch.index_select( + w2_sf_u8, + 1, + perm_w2, + out=_mxfp8_scratch("w2_sf", w2_sf_u8.shape, w2_sf_u8.device), + ) + w13_scale_shuffled = ( + block_scale_interleave(w13_sf_gathered) + .view(MXFP8_SCALE_DTYPE) + .view(num_experts, -1) + ) + w2_scale_shuffled = ( + block_scale_interleave(w2_sf_gathered) + .view(MXFP8_SCALE_DTYPE) + .view(num_experts, -1) + ) + return ( + w13_shuffled.view(MXFP8_VALUE_DTYPE), + w2_shuffled.view(MXFP8_VALUE_DTYPE), + w13_scale_shuffled, + w2_scale_shuffled, + ) + + +def _shuffle_mxfp8_moe_per_expert( + w13_weight: torch.Tensor, + w2_weight: torch.Tensor, + w13_scale: torch.Tensor, + w2_scale: torch.Tensor, + is_gated: bool, + epilogue_tile_m: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Per-expert reference shuffle, kept for NRL_MXFP8_SHUFFLE_VERIFY.""" from flashinfer import ( reorder_rows_for_gated_act_gemm, shuffle_matrix_a, shuffle_matrix_sf_a, ) - from vllm.model_executor.layers.fused_moe.layer import FusedMoeWeightScaleSupported - from vllm.model_executor.layers.quantization.utils.flashinfer_utils import ( - swap_w13_to_w31, - ) from vllm.model_executor.layers.quantization.utils.mxfp8_utils import ( MXFP8_SCALE_DTYPE, MXFP8_VALUE_DTYPE, ) - from vllm.model_executor.parameter import ModelWeightParameter - from vllm.model_executor.utils import set_weight_attrs - - epilogue_tile_m = 128 - num_experts = layer.w13_weight.shape[0] - is_gated = self.moe.is_act_and_mul - intermediate_size_factor = 2 if is_gated else 1 - - w13_weight = layer.w13_weight.data - if not hasattr(layer, "w13_weight_scale_from_checkpoint"): - w13_scale = layer.w13_weight_scale.data - else: - w13_scale = layer.w13_weight_scale_from_checkpoint.data - if is_gated: - # FI TRTLLM gated kernels use W31 ordering. Model checkpoints store - # gated projection as W13, so convert once before shuffling. - w13_weight = swap_w13_to_w31(w13_weight) - w13_scale = swap_w13_to_w31(w13_scale) - w2_weight = layer.w2_weight.data - if not hasattr(layer, "w2_weight_scale_from_checkpoint"): - w2_scale = layer.w2_weight_scale.data - else: - w2_scale = layer.w2_weight_scale_from_checkpoint.data + num_experts = w13_weight.shape[0] + w13_rows = w13_weight.shape[1] + w2_rows = w2_weight.shape[1] w13_weight_shuffled = [] w2_weight_shuffled = [] w13_scale_shuffled = [] w2_scale_shuffled = [] for i in range(num_experts): - w13_i = w13_weight[i].reshape( - intermediate_size_factor * layer.intermediate_size_per_partition, -1 - ) - w13_sf_i = w13_scale[i].reshape( - intermediate_size_factor * layer.intermediate_size_per_partition, -1 - ) + w13_i = w13_weight[i].reshape(w13_rows, -1) + w13_sf_i = w13_scale[i].reshape(w13_rows, -1) if is_gated: # Reorder rows for gated activation layout expected by TRTLLM. w13_i = reorder_rows_for_gated_act_gemm(w13_i.clone()) @@ -903,18 +1089,11 @@ def process_weights_after_loading_mxfp8_moe(self, layer) -> None: w13_weight_shuffled.append(w13_shuffled_i.contiguous().view(MXFP8_VALUE_DTYPE)) w2_weight_shuffled.append(w2_shuffled_i.contiguous().view(MXFP8_VALUE_DTYPE)) w13_sf_shuffled_i = shuffle_matrix_sf_a( - pad_flashinfer_scale_k( - w13_sf_i.view(torch.uint8).reshape( - intermediate_size_factor * layer.intermediate_size_per_partition, - -1, - ) - ), + pad_flashinfer_scale_k(w13_sf_i.view(torch.uint8).reshape(w13_rows, -1)), epilogue_tile_m, ) w2_sf_shuffled_i = shuffle_matrix_sf_a( - pad_flashinfer_scale_k( - w2_scale[i].view(torch.uint8).reshape(layer.hidden_size, -1) - ), + pad_flashinfer_scale_k(w2_scale[i].view(torch.uint8).reshape(w2_rows, -1)), epilogue_tile_m, ) w13_scale_shuffled.append( @@ -922,7 +1101,139 @@ def process_weights_after_loading_mxfp8_moe(self, layer) -> None: ) w2_scale_shuffled.append(w2_sf_shuffled_i.contiguous().view(MXFP8_SCALE_DTYPE)) - if not hasattr(layer, "w13_weight_scale_from_checkpoint"): + return ( + torch.stack(w13_weight_shuffled).contiguous(), + torch.stack(w2_weight_shuffled).contiguous(), + torch.stack(w13_scale_shuffled).contiguous(), + torch.stack(w2_scale_shuffled).contiguous(), + ) + + +def process_weights_after_loading_mxfp8_moe(self, layer) -> None: + """Shuffle weights and scales into FlashInfer TRTLLM MXFP8 layout. + + FlashInfer's shuffle_matrix_sf_a requires M divisible by 128, so the + intermediate dim is padded to 128 (Nemotron-3-Nano: 1856 -> 1920, + 928 -> 1024 at TP2). The TRTLLM kernel also produced non-finite output + for Nano's hidden 2816 unless padded to the next 512 boundary (3072). + When padding is needed, the checkpoint-layout parameters keep their + original shapes so refit keeps loading into them unchanged, and the + padded+shuffled kernel inputs are maintained as separate *_for_apply + tensors consumed by apply_monolithic_mxfp8_moe. Already-aligned models + keep the original in-place shuffle behavior. + """ + from vllm.model_executor.layers.fused_moe.layer import FusedMoeWeightScaleSupported + from vllm.model_executor.layers.quantization.utils.flashinfer_utils import ( + swap_w13_to_w31, + ) + from vllm.model_executor.layers.quantization.utils.mxfp8_utils import ( + MXFP8_BLOCK_SIZE, + ) + from vllm.model_executor.parameter import ModelWeightParameter + from vllm.model_executor.utils import set_weight_attrs + + epilogue_tile_m = 128 + is_gated = self.moe.is_act_and_mul + intermediate_size_factor = 2 if is_gated else 1 + + # Sizes come from the checkpoint-layout params, not layer attributes: + # layer.intermediate_size_per_partition is switched to the padded value + # below, while these params keep their original shapes across refits. + original_intermediate_size = layer.w2_weight.shape[2] + original_hidden_size = layer.w13_weight.shape[2] + padded_intermediate_size = _round_up(original_intermediate_size, 128) + padded_hidden_size = _round_up(original_hidden_size, 512) + needs_padding = ( + padded_intermediate_size != original_intermediate_size + or padded_hidden_size != original_hidden_size + ) + + first_load = not hasattr(layer, "w13_weight_scale_from_checkpoint") + w13_weight = layer.w13_weight.data + if first_load: + w13_scale = layer.w13_weight_scale.data + w2_scale = layer.w2_weight_scale.data + else: + w13_scale = layer.w13_weight_scale_from_checkpoint.data + w2_scale = layer.w2_weight_scale_from_checkpoint.data + w2_weight = layer.w2_weight.data + + if needs_padding: + w13_weight = _pad_w13_shards( + w13_weight, intermediate_size_factor, padded_intermediate_size + ) + w13_weight = _pad_tensor_dim(w13_weight, 2, padded_hidden_size) + w13_scale = _pad_w13_shards( + w13_scale, intermediate_size_factor, padded_intermediate_size, pad_value=1 + ) + w13_scale = _pad_tensor_dim( + w13_scale, 2, padded_hidden_size // MXFP8_BLOCK_SIZE, pad_value=1 + ) + w2_weight = _pad_tensor_dim(w2_weight, 1, padded_hidden_size) + w2_weight = _pad_tensor_dim(w2_weight, 2, padded_intermediate_size) + w2_scale = _pad_tensor_dim(w2_scale, 1, padded_hidden_size, pad_value=1) + w2_scale = _pad_tensor_dim( + w2_scale, 2, padded_intermediate_size // MXFP8_BLOCK_SIZE, pad_value=1 + ) + # Zero E8M0 scale bytes destabilize the TRTLLM kernel; clamp to byte 1. + w13_scale = _clamp_mxfp8_scale(w13_scale) + w2_scale = _clamp_mxfp8_scale(w2_scale) + + if is_gated: + # FI TRTLLM gated kernels use W31 ordering. Model checkpoints store + # gated projection as W13, so convert once before shuffling. + w13_weight = swap_w13_to_w31(w13_weight) + w13_scale = swap_w13_to_w31(w13_scale) + + # NRL_MXFP8_BATCHED_SHUFFLE=0 is the kill switch back to the per-expert path. + use_batched_shuffle = os.getenv("NRL_MXFP8_BATCHED_SHUFFLE", "1") != "0" + if use_batched_shuffle: + ( + w13_weight_shuffled, + w2_weight_shuffled, + w13_scale_shuffled, + w2_scale_shuffled, + ) = _shuffle_mxfp8_moe_batched( + layer, w13_weight, w2_weight, w13_scale, w2_scale, is_gated, epilogue_tile_m + ) + else: + ( + w13_weight_shuffled, + w2_weight_shuffled, + w13_scale_shuffled, + w2_scale_shuffled, + ) = _shuffle_mxfp8_moe_per_expert( + w13_weight, w2_weight, w13_scale, w2_scale, is_gated, epilogue_tile_m + ) + + global mxfp8_shuffle_verified + if ( + use_batched_shuffle + and os.getenv("NRL_MXFP8_SHUFFLE_VERIFY") == "1" + and not mxfp8_shuffle_verified + ): + reference = _shuffle_mxfp8_moe_per_expert( + w13_weight, w2_weight, w13_scale, w2_scale, is_gated, epilogue_tile_m + ) + batched = ( + w13_weight_shuffled, + w2_weight_shuffled, + w13_scale_shuffled, + w2_scale_shuffled, + ) + for got, want, tensor_name in zip( + batched, reference, ("w13_weight", "w2_weight", "w13_scale", "w2_scale") + ): + assert torch.equal(got.view(torch.uint8), want.view(torch.uint8)), ( + f"Batched MXFP8 shuffle mismatch vs per-expert reference: {tensor_name}" + ) + mxfp8_shuffle_verified = True + print( + "[NRL_MXFP8_SHUFFLE_VERIFY] batched MoE shuffle matches the " + "per-expert reference bit-exactly" + ) + + if first_load: layer.w13_weight_scale_from_checkpoint = ModelWeightParameter( data=layer.w13_weight_scale.data, input_dim=2, @@ -955,17 +1266,146 @@ def process_weights_after_loading_mxfp8_moe(self, layer) -> None: layer.w2_weight_scale_from_checkpoint, {"quant_method": FusedMoeWeightScaleSupported.BLOCK.value}, ) - layer.w13_weight_scale = torch.nn.Parameter( - torch.stack(w13_scale_shuffled).contiguous(), requires_grad=False + + if needs_padding: + # Checkpoint-layout params (w13_weight, w2_weight, w13_weight_scale, + # w2_weight_scale and the *_from_checkpoint aliases) stay untouched so + # refit keeps loading unpadded tensors into them. The kernel consumes + # the *_for_apply tensors instead. + layer.mxfp8_unpadded_hidden_size = original_hidden_size + layer.mxfp8_padded_hidden_size = padded_hidden_size + layer.mxfp8_unpadded_intermediate_size_per_partition = ( + original_intermediate_size ) - layer.w2_weight_scale = torch.nn.Parameter( - torch.stack(w2_scale_shuffled).contiguous(), requires_grad=False + layer.mxfp8_padded_intermediate_size_per_partition = padded_intermediate_size + # FusedMoE.intermediate_size_per_partition is a read-only property + # backed by moe_config; the TRTLLM kernel must see the padded value. + # hidden stays original at the layer level: apply pads x and narrows + # the output back. + layer.moe_config.intermediate_size_per_partition = padded_intermediate_size + self.moe.intermediate_size_per_partition = padded_intermediate_size + _set_mxfp8_apply_tensor(layer, "w13_weight_for_apply", w13_weight_shuffled) + _set_mxfp8_apply_tensor(layer, "w2_weight_for_apply", w2_weight_shuffled) + _set_mxfp8_apply_tensor(layer, "w13_scale_for_apply", w13_scale_shuffled) + _set_mxfp8_apply_tensor(layer, "w2_scale_for_apply", w2_scale_shuffled) + else: + if first_load: + layer.w13_weight_scale = torch.nn.Parameter( + w13_scale_shuffled, requires_grad=False + ) + layer.w2_weight_scale = torch.nn.Parameter( + w2_scale_shuffled, requires_grad=False + ) + else: + layer.w13_weight_scale.copy_(w13_scale_shuffled) + layer.w2_weight_scale.copy_(w2_scale_shuffled) + layer.w13_weight.copy_(w13_weight_shuffled) + layer.w2_weight.copy_(w2_weight_shuffled) + + +def apply_monolithic_mxfp8_moe( + self, + layer, + x: torch.Tensor, + router_logits: torch.Tensor, + input_ids: torch.Tensor | None = None, +) -> torch.Tensor: + """Forward for the FlashInfer TRTLLM MXFP8 MoE with hidden-dim padding. + + Mirrors vLLM 0.20.0 ModelOptMxFp8FusedMoE.apply_monolithic, with three + changes: reads the *_for_apply tensors built by + process_weights_after_loading_mxfp8_moe when padding is active, pads x's + hidden dim to mxfp8_padded_hidden_size before the kernel and narrows the + output back, and allows RELU2_NO_MUL for non-gated MoEs (Nemotron-3-Nano). + """ + from flashinfer.fused_moe.core import ( + ActivationType, + Fp8QuantizationType, + ) + from vllm.model_executor.layers.fused_moe.activation import MoEActivation + from vllm.model_executor.layers.fused_moe.config import RoutingMethodType + from vllm.model_executor.layers.fused_moe.oracle.fp8 import Fp8MoeBackend + from vllm.model_executor.layers.quantization.utils.mxfp8_utils import ( + mxfp8_e4m3_quantize, + ) + from vllm.utils.flashinfer import flashinfer_trtllm_fp8_block_scale_moe + + assert self.mxfp8_backend == Fp8MoeBackend.FLASHINFER_TRTLLM + + if layer.enable_eplb: + raise NotImplementedError( + "EPLB is not supported for FlashInfer TRTLLM MXFP8 MoE backend." + ) + + # Map vLLM MoEActivation to FlashInfer ActivationType. + activation_map = { + MoEActivation.SILU: ActivationType.Swiglu, + MoEActivation.RELU2_NO_MUL: ActivationType.Relu2, + } + if layer.activation not in activation_map: + raise NotImplementedError( + "FlashInfer TRTLLM MXFP8 MoE supports only " + f"{list(activation_map)}, got {layer.activation}." + ) + fi_activation_type = activation_map[layer.activation] + + # DeepSeekV3 routing requires float32 logits; others expect bfloat16. + if layer.routing_method_type == RoutingMethodType.DeepSeekV3: + assert router_logits.dtype == torch.float32, ( + "DeepSeekV3 routing requires float32 router_logits, " + f"got {router_logits.dtype}." ) else: - layer.w13_weight_scale.copy_(torch.stack(w13_scale_shuffled).contiguous()) - layer.w2_weight_scale.copy_(torch.stack(w2_scale_shuffled).contiguous()) - layer.w13_weight.copy_(torch.stack(w13_weight_shuffled).contiguous()) - layer.w2_weight.copy_(torch.stack(w2_weight_shuffled).contiguous()) + router_logits = router_logits.to(torch.bfloat16) + + # Treat 0 as "unset" for compatibility with ungrouped routing configs. + n_group = layer.num_expert_group or None + topk_group = layer.topk_group or None + + unpadded_hidden_size = x.shape[-1] + padded_hidden_size = getattr( + layer, "mxfp8_padded_hidden_size", unpadded_hidden_size + ) + if unpadded_hidden_size < padded_hidden_size: + x = torch.nn.functional.pad( + x, (0, padded_hidden_size - unpadded_hidden_size), value=0.0 + ) + + hidden_states_mxfp8, hidden_states_scale = mxfp8_e4m3_quantize( + x, + is_sf_swizzled_layout=False, + ) + + output = flashinfer_trtllm_fp8_block_scale_moe( + routing_logits=router_logits, + routing_bias=layer.e_score_correction_bias, + hidden_states=hidden_states_mxfp8, + hidden_states_scale=hidden_states_scale, + gemm1_weights=getattr(layer, "w13_weight_for_apply", layer.w13_weight), + gemm1_weights_scale=getattr( + layer, "w13_scale_for_apply", layer.w13_weight_scale + ), + gemm2_weights=getattr(layer, "w2_weight_for_apply", layer.w2_weight), + gemm2_weights_scale=getattr(layer, "w2_scale_for_apply", layer.w2_weight_scale), + num_experts=layer.global_num_experts, + top_k=layer.top_k, + # Keep Optional semantics: FlashInfer expects None for non-grouped + # routing (e.g. Qwen3 Renormalize), not 0. + n_group=n_group, + topk_group=topk_group, + intermediate_size=layer.intermediate_size_per_partition, + local_expert_offset=layer.ep_rank * layer.local_num_experts, + local_num_experts=layer.local_num_experts, + routed_scaling_factor=layer.routed_scaling_factor, + routing_method_type=layer.routing_method_type, + use_shuffled_weight=True, + weight_layout=0, + fp8_quantization_type=Fp8QuantizationType.MxFp8, + activation_type=fi_activation_type, + ) + if output.shape[-1] != unpadded_hidden_size: + output = output[..., :unpadded_hidden_size].contiguous() + return output def process_weights_after_loading_kv(self, layer) -> None: diff --git a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py index ac4db666cfe..e2cea680595 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py @@ -13,6 +13,88 @@ # limitations under the License. +import torch + +MXFP8_BLOCK_SIZE = 32 +MXFP8_VALUE_DTYPE = torch.float8_e4m3fn + + +def _mxfp8_e4m3_quantize_torch( + x: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + """Reference MXFP8 quantization with row-major scales. + + Replicates vLLM's _mxfp8_e4m3_quantize_torch (Apache-2.0, + vllm/model_executor/layers/quantization/utils/mxfp8_utils.py) for trainer + processes without a vLLM install: for each block of 32 elements along the + last dimension, a shared e8m0 scale (biased exponent of the block amax) + and float8_e4m3fn values. + """ + assert x.shape[-1] % MXFP8_BLOCK_SIZE == 0, ( + f"MXFP8 requires the last dim to be divisible by {MXFP8_BLOCK_SIZE}, got {x.shape}" + ) + orig_shape = x.shape + num_blocks = x.shape[-1] // MXFP8_BLOCK_SIZE + + x_fp32 = x.to(torch.float32) + x_blocked = x_fp32.view(*orig_shape[:-1], num_blocks, MXFP8_BLOCK_SIZE) + + amax = x_blocked.abs().amax(dim=-1) + amax = amax.clamp(min=torch.finfo(torch.float32).tiny) + scale_biased = torch.floor(torch.log2(amax)) + 127.0 + scale_biased = scale_biased.clamp(0, 254) + scales_uint8 = scale_biased.to(torch.uint8) + + descale = torch.exp2(scale_biased - 127.0) + x_scaled = x_blocked / descale.unsqueeze(-1) + + x_fp8 = x_scaled.view(orig_shape).to(MXFP8_VALUE_DTYPE) + + if x.ndim == 2: + scales_uint8 = scales_uint8.view(x.shape[0], -1) + elif x.ndim == 3: + scales_uint8 = scales_uint8.view(x.shape[0], x.shape[1], -1) + + return x_fp8, scales_uint8 + + +def mxfp8_e4m3_quantize_for_refit( + x: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + """Quantize a weight to MXFP8 on the trainer for pre-quantized refit. + + Mirrors the receiver path in quantization/fp8.py load_weights + (mxfp8_e4m3_quantize + scale squeeze) so the streamed E4M3 data and + *_scale_from_checkpoint scales load bit-identically without receiver-side + re-quantization. Uses the same flashinfer kernel as vLLM on Blackwell and + the torch reference elsewhere. + """ + x_q = x_scales = None + # Kernel dispatch keys off the TRAINER GPU while the receiver keys off the + # inference GPU. On homogeneous clusters both take the same path; on mixed + # Hopper/Blackwell clusters the flashinfer and torch paths may differ in + # boundary rounding - validate with the parity test before relying on it. + if x.is_cuda and torch.cuda.get_device_capability(x.device) >= (10, 0): + try: + from flashinfer import mxfp8_quantize as flashinfer_mxfp8_quantize + except ImportError: + pass + else: + x_q, x_scales = flashinfer_mxfp8_quantize( + x, is_sf_swizzled_layout=False, alignment=32 + ) + if x_scales.ndim == 1 and x.ndim == 2: + x_scales = x_scales.view(x.size(0), -1) + if x_q is None or x_scales is None: + x_q, x_scales = _mxfp8_e4m3_quantize_torch(x) + x_scales = torch.squeeze(x_scales, dim=-1) + # Match the receiver path's zero-scale clamp: an E8M0 byte of 0 (2^-127) + # destabilizes the TRTLLM kernels, and pre-quantized tensors skip the + # receiver-side quantize branch where the clamp normally runs. + x_scales = torch.where(x_scales == 0, torch.ones_like(x_scales), x_scales) + return x_q, x_scales + + def get_vllm_qkv_scale_names(layer_idx: int) -> dict[str, str]: """Get vLLM-compatible parameter names for Q/K/V FP8 scales. diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index d21f169a261..4c0e201d8c3 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -12,12 +12,13 @@ # See the License for the specific language governing permissions and # limitations under the License. import gc +import os import re import socket import traceback from collections.abc import Callable, Iterable, Iterator, Sequence from contextlib import contextmanager -from typing import Any, Literal +from typing import Any, Literal, Optional import torch import zmq @@ -47,7 +48,7 @@ ) -WeightUpdateTransport = Literal["ipc", "collective"] +WeightUpdateTransport = Literal["ipc", "collective", "checkpoint_engine"] WeightUpdateFinalizer = Callable[[], None] @@ -126,6 +127,138 @@ def fix_gemma3_vision_weight_name(key: str) -> str: ) +class _RefitLoaderCache: + """Recorded weight_loader calls for refit weight names. + + vLLM's model.load_weights re-resolves every weight name through the + model's stacked/expert parameter mappings on each call; for large MoE + refits that is millions of substring checks per refit over a key set that + is static after the prepare_refit_info handshake. This cache records, per + name, the (loader, param, args, kwargs) of every weight_loader call the + first time a name is loaded and replays them directly afterwards. + """ + + def __init__(self) -> None: + self.calls: dict[str, list[tuple[Any, torch.nn.Parameter, tuple, dict]]] = {} + # Names whose loads never reached a wrapped weight_loader (skipped, + # transformed before dispatch, or default-loaded); these keep going + # through model.load_weights. + self.uncached: set[str] = set() + self.snapshot: dict[str, torch.nn.Parameter] = {} + + def reset(self) -> None: + self.calls.clear() + self.uncached.clear() + self.snapshot.clear() + + +def _cached_params_still_valid(model: Any, cache: _RefitLoaderCache) -> bool: + current = dict(model.named_parameters()) + return all(current.get(name) is param for name, param in cache.snapshot.items()) + + +def _record_loader_calls( + model: Any, cache: _RefitLoaderCache, weights: list[tuple[str, torch.Tensor]] +) -> set[str]: + """Run model.load_weights once while recording every weight_loader call. + + Incoming weights are matched to loader calls by tensor object identity, + so only loads that pass the original tensor through a parameter's + weight_loader attribute are captured; everything else lands in + cache.uncached. Returns the loaded names from model.load_weights. + """ + from vllm.model_executor.model_loader.weight_utils import default_weight_loader + + weight_names = {id(weight): name for name, weight in weights} + recorded: dict[str, list] = {} + originals: list[tuple[torch.nn.Parameter, Any]] = [] + + def make_recorder(loader): + def recorder(param, loaded_weight, *args, **kwargs): + name = weight_names.get(id(loaded_weight)) + if name is not None: + recorded.setdefault(name, []).append((loader, param, args, kwargs)) + return loader(param, loaded_weight, *args, **kwargs) + + return recorder + + try: + for param_name, param in model.named_parameters(): + loader = getattr(param, "weight_loader", None) + # Leave default_weight_loader params unwrapped: model + # load_weights implementations dispatch on + # `weight_loader == default_weight_loader` with a different + # argument list. + if loader is None or loader is default_weight_loader: + continue + cache.snapshot[param_name] = param + originals.append((param, loader)) + param.weight_loader = make_recorder(loader) + loaded = model.load_weights(weights) + finally: + for param, loader in originals: + param.weight_loader = loader + + for name, _ in weights: + calls = recorded.get(name) + if calls is None: + cache.uncached.add(name) + else: + cache.calls[name] = calls + return loaded if loaded is not None else set() + + +def load_weights_maybe_cached( + model: Any, weights: list[tuple[str, torch.Tensor]] +) -> set[str]: + """model.load_weights with optional loader replay caching. + + Opt-in via NRL_REFIT_CACHED_LOADERS=1 since model load_weights + implementations vary; the default is a plain model.load_weights call. + Cached parameter identities are re-validated against named_parameters() + on every call, so a process_weights_after_loading pass that replaces + parameter objects drops the cache instead of loading into orphans. + Returns the set of loaded weight names, mirroring model.load_weights. + """ + if os.getenv("NRL_REFIT_CACHED_LOADERS") != "1": + return model.load_weights(weights) + + cache = getattr(model, "_nrl_refit_loader_cache", None) + if cache is None: + cache = _RefitLoaderCache() + model._nrl_refit_loader_cache = cache + + replay = [] + fallback = [] + record = [] + for name, weight in weights: + if name in cache.calls: + replay.append((name, weight)) + elif name in cache.uncached: + fallback.append((name, weight)) + else: + record.append((name, weight)) + + if replay and not _cached_params_still_valid(model, cache): + cache.reset() + return model.load_weights(weights) + + loaded: set[str] = set() + for name, weight in replay: + for loader, param, args, kwargs in cache.calls[name]: + # Expert loaders return False for non-local shards; a name only + # counts as loaded when some call does not report failure. + if loader(param, weight, *args, **kwargs) is not False: + loaded.add(name) + if record: + loaded |= _record_loader_calls(model, cache, record) + if fallback: + fallback_loaded = model.load_weights(fallback) + if fallback_loaded is not None: + loaded |= fallback_loaded + return loaded + + def _read_mtp_layer_weights_from_checkpoint( model_path: str, mtp_layer_indices: set[int] ) -> list[tuple[str, torch.Tensor]]: @@ -144,7 +277,6 @@ def _read_mtp_layer_weights_from_checkpoint( tensors on CPU. """ import json - import os from safetensors import safe_open @@ -183,7 +315,9 @@ def _get_named_parameters(self) -> dict[str, torch.nn.Parameter]: def _load_full_hf_weights( self, policy_weights: list[tuple[str, torch.Tensor]] ) -> None: - self.model_runner.model.load_weights(weights=policy_weights) + # Refit optimization: replay cached weight-loader routing when + # NRL_REFIT_CACHED_LOADERS=1 (identity-validated), else plain load. + load_weights_maybe_cached(self.model_runner.model, policy_weights) def _load_hf_weights(self, policy_weights: list[tuple[str, torch.Tensor]]) -> None: from nemo_rl.models.generation.vllm.quantization import fp8 @@ -263,15 +397,39 @@ def maybe_init_zmq(self): self.zmq_socket.setsockopt(zmq.LINGER, 0) self.zmq_socket.connect(self.get_zmq_address()) - def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: + def prepare_refit_info( + self, state_dict_info: dict[str, Any] + ) -> Optional[list[str]]: """Prepare state dict metadata for weight refitting and IPC streaming. Args: state_dict_info (dict): A dictionary containing the info for refit. e.g. {tensor_name: (shape, dtype)} + + Returns: + When MXFP8 trainer-side pre-quantization is enabled + (vllm_cfg.refit_prequantize), the list of parameter names this + worker will quantize at load time; the trainer quantizes exactly + these and streams E4M3 data plus *_scale_from_checkpoint scales. + None otherwise. """ self.state_dict_info = state_dict_info # pyrefly: ignore[implicitly-defined-attribute] This class does not define __init__ so assignments like this should be ignored + from nemo_rl.models.generation.vllm.quantization import fp8 + + if not ( + fp8.global_fp8_config is not None + and fp8.global_fp8_config.is_mx + and fp8.global_fp8_config.refit_prequantize + and fp8.is_fp8_model(self.model_runner.vllm_config) + ): + return None + return [ + name + for name in state_dict_info + if fp8._is_fp8_weight(name, self.model_runner.model) + ] + def prepare_sparse_delta_refit_info( self, state_dict_info: dict[str, tuple[tuple[int, ...], torch.dtype]] ) -> list[str]: @@ -286,28 +444,6 @@ def _uses_fp8_kv_cache(self) -> bool: kv_cache_dtype = getattr(cache_config, "cache_dtype", None) return kv_cache_dtype is not None and "fp8" in str(kv_cache_dtype).lower() - def _maybe_process_fp8_kv_cache(self) -> None: - """Process weights after loading for FP8 KV cache (static scales).""" - if not self._uses_fp8_kv_cache(): - return - - # FP8 KV cache: process KV scales after weight loading - from vllm.config import set_current_vllm_config - from vllm.model_executor.model_loader.utils import ( - process_weights_after_loading, - ) - - # Get target device for processing - target_device = next(self.model_runner.model.parameters()).device - - # Call process_weights_after_loading to handle KV scales - with set_current_vllm_config(self.model_runner.vllm_config): - process_weights_after_loading( - self.model_runner.model, - self.model_runner.model_config, - target_device, - ) - @staticmethod def _split_policy_and_draft_weights( weights: list[tuple[str, torch.Tensor]], @@ -489,9 +625,8 @@ def finalize() -> None: ) yield finalize - # Preserve the IPC lifetime boundary: the COMPLETE ACK is sent before - # this optional second pass, just as it was before lifecycle hooks. - self._maybe_process_fp8_kv_cache() + # KV-cache scales are covered by the full process_weights_after_loading + # pass in finalize(); no second pass is needed. def _weight_update_errors_are_fatal(self) -> bool: """Whether transport errors should propagate instead of returning False.""" diff --git a/nemo_rl/models/generation/vllm/vllm_generation.py b/nemo_rl/models/generation/vllm/vllm_generation.py index e889fe9ed2d..e760edc6e55 100644 --- a/nemo_rl/models/generation/vllm/vllm_generation.py +++ b/nemo_rl/models/generation/vllm/vllm_generation.py @@ -913,8 +913,16 @@ def shutdown(self) -> bool: print(f"Error during policy shutdown: {e}") return False - def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: - """Prepare the info for refit.""" + def prepare_refit_info( + self, state_dict_info: Optional[dict[str, Any]] + ) -> Optional[list[str]]: + """Prepare the info for refit. + + Returns: + When MXFP8 trainer-side pre-quantization is enabled + (vllm_cfg.refit_prequantize), the parameter names the engine wants + quantized on the trainer before streaming. None otherwise. + """ # Choose the appropriate method based on async_engine setting method_name = ( "prepare_refit_info_async" @@ -929,8 +937,13 @@ def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], ) - # Wait for all futures to complete - ray.get(futures) + # Union the fp8-eligible parameter names across workers: replicas are + # equivalent, but with pipeline parallelism each worker only reports + # the parameters of its local shard. + names = sorted( + {name for result in ray.get(futures) if result for name in result} + ) + return names or None def update_weights_via_ipc_zmq(self) -> list[ray.ObjectRef]: """Update weights of the policy using IPC handles via ZMQ socket.""" diff --git a/nemo_rl/models/generation/vllm/vllm_worker.py b/nemo_rl/models/generation/vllm/vllm_worker.py index a5063d8d43f..62a4cb083c0 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker.py +++ b/nemo_rl/models/generation/vllm/vllm_worker.py @@ -1029,9 +1029,19 @@ def report_device_id(self) -> list[str]: ) return cast(list[str], list_of_worker_results) - def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: - """Prepare the info for refit.""" - self.llm.collective_rpc("prepare_refit_info", args=(state_dict_info,)) + def prepare_refit_info( + self, state_dict_info: dict[str, Any] + ) -> Optional[list[str]]: + """Prepare the info for refit. + + Returns the parameter names the engine wants pre-quantized on the + trainer (vllm_cfg.refit_prequantize), or None. + """ + results = self.llm.collective_rpc("prepare_refit_info", args=(state_dict_info,)) + # Union across the engine's TP/PP workers: with pipeline parallelism + # each shard only classifies its local parameters as fp8-eligible. + names = sorted({name for result in results if result for name in result}) + return names or None @wrap_with_nvtx_name("vllm_genertion_worker/update_weights_via_ipc_zmq") def update_weights_via_ipc_zmq(self) -> bool: diff --git a/nemo_rl/models/generation/vllm/vllm_worker_async.py b/nemo_rl/models/generation/vllm/vllm_worker_async.py index b08184196b9..d6f0a8214fc 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker_async.py +++ b/nemo_rl/models/generation/vllm/vllm_worker_async.py @@ -1405,9 +1405,17 @@ async def report_device_id_async(self) -> list[str]: return cast(list[str], list_of_worker_results) - async def prepare_refit_info_async(self, state_dict_info: dict[str, Any]) -> None: + async def prepare_refit_info_async( + self, state_dict_info: dict[str, Any] + ) -> Optional[list[str]]: """Async version of prepare_refit_info.""" - await self.llm.collective_rpc("prepare_refit_info", args=(state_dict_info,)) + results = await self.llm.collective_rpc( + "prepare_refit_info", args=(state_dict_info,) + ) + # Union across the engine's TP/PP workers: with pipeline parallelism + # each shard only classifies its local parameters as fp8-eligible. + names = sorted({name for result in results if result for name in result}) + return names or None async def update_weights_via_ipc_zmq_async( self, diff --git a/nemo_rl/models/megatron/setup.py b/nemo_rl/models/megatron/setup.py index d5aaae2d6fc..43ee7bd5b20 100644 --- a/nemo_rl/models/megatron/setup.py +++ b/nemo_rl/models/megatron/setup.py @@ -1738,6 +1738,10 @@ def composed_peft_hook(model: list[MegatronModule]) -> list[MegatronModule]: ) reference_state_dict = {} + # NotRequired key: absent means disabled, default lives in the exemplar YAML. + pinned_reference_swap = bool( + config["megatron_cfg"].get("pinned_reference_swap") + ) if should_load_checkpoint or use_peft: reference_model = reference_model[0] @@ -1745,13 +1749,23 @@ def composed_peft_hook(model: list[MegatronModule]) -> list[MegatronModule]: # Store reference state dict on CPU for name, item in reference_model.state_dict().items(): if isinstance(item, torch.Tensor): - cpu_item = item.detach().to( - device="cpu", non_blocking=True, copy=True - ) + if pinned_reference_swap: + # Pinned so use_reference_model can upload the reference + # weights with non-blocking H2D copies each step. + cpu_item = torch.empty( + item.shape, dtype=item.dtype, device="cpu", pin_memory=True + ) + cpu_item.copy_(item.detach(), non_blocking=True) + else: + cpu_item = item.detach().to( + device="cpu", non_blocking=True, copy=True + ) del item else: cpu_item = item reference_state_dict[name] = cpu_item + if pinned_reference_swap: + torch.cuda.synchronize() print("Reference model loaded") else: print("Reference model not loaded") diff --git a/nemo_rl/models/policy/__init__.py b/nemo_rl/models/policy/__init__.py index fcda0190e36..80cab60b3ec 100644 --- a/nemo_rl/models/policy/__init__.py +++ b/nemo_rl/models/policy/__init__.py @@ -328,6 +328,17 @@ class MegatronConfig(TypedDict): # 1 is the minimum recommendation for RL since we almost always need to offload before beginning generation. # Setting to 0 is faster, but you are more likely to run out of GPU memory. In SFT/DPO, the default is 0. empty_unused_memory_level: int + # When True, offload_after_refit skips rerunning the full offload_before_refit + # pass (grad-buffer moves, cache clears, and a second gc.collect/empty_cache) + # and only re-offloads the optimizer with a single allocator cleanup. + refit_slim_offload_after: NotRequired[bool] + # When True, the reference-policy swap in use_reference_model stages the + # active weights in persistent pinned CPU buffers and keeps the reference + # state dict pinned from setup, so the three per-step full-model PCIe + # transfers run as non-blocking pinned copies instead of pageable ones. + # Costs 2x model-size pinned host RAM per worker. Default False keeps the + # current pageable-copy behavior. + pinned_reference_swap: NotRequired[bool] activation_checkpointing: bool # Recompute granularity: "full" recomputes all activations, "selective" recomputes # only specific modules (see recompute_modules). "selective" typically saves ~10-18GB @@ -540,6 +551,11 @@ class PolicyConfig(TypedDict): # This sets the clipping norm for the DTensorPolicyWorkers (Megatron's is called clip_grad) max_grad_norm: NotRequired[float | int | None] refit_buffer_size_gb: NotRequired[float] + # Keep the CUDA-IPC ping-pong staging buffers allocated across refits instead of + # reallocating and freeing them (plus a gc.collect/empty_cache pair) every refit. + # Works best with a fixed refit_buffer_size_gb; the buffers stay resident on the + # trainer GPU between refits (2x half-buffer bytes). + refit_persistent_ipc_buffers: NotRequired[bool] optimizer: NotRequired[PytorchOptimizerConfig | None] scheduler: NotRequired[ list[SinglePytorchSchedulerConfig | SinglePytorchMilestonesConfig] diff --git a/nemo_rl/models/policy/interfaces.py b/nemo_rl/models/policy/interfaces.py index f0c1ad6bb8e..865acd9f31b 100644 --- a/nemo_rl/models/policy/interfaces.py +++ b/nemo_rl/models/policy/interfaces.py @@ -185,6 +185,18 @@ def offload_to_cpu(self) -> None: def prepare_refit_info(self) -> Optional[dict[str, Any]]: pass + def enable_refit_prequantize( + self, param_names: list[str] + ) -> Optional[dict[str, Any]]: + """Quantize the listed params on the trainer during refit streaming. + + Returns: + Refit info updated with the quantized dtypes and scale entries. + """ + raise NotImplementedError( + "enable_refit_prequantize is not implemented for this policy worker" + ) + @abstractmethod def stream_weights_via_ipc_zmq( self, *args: Any, **kwargs: Any diff --git a/nemo_rl/models/policy/lm_policy.py b/nemo_rl/models/policy/lm_policy.py index 397b4e086b5..7302562805a 100644 --- a/nemo_rl/models/policy/lm_policy.py +++ b/nemo_rl/models/policy/lm_policy.py @@ -940,6 +940,20 @@ def prepare_refit_info(self) -> Optional[dict[str, Any]]: # Only get the first worker's info since all workers will have the same result return results[0] + def enable_refit_prequantize( + self, param_names: list[str] + ) -> Optional[dict[str, Any]]: + """Enable trainer-side MXFP8 quantization of the listed params for refit. + + Returns: + dict: Refit info updated with quantized dtypes and scale entries. + """ + futures = self.worker_group.run_all_workers_single_data( + "enable_refit_prequantize", param_names=param_names + ) + results = ray.get(futures) + return results[0] + def finish_inference(self) -> None: """Offload policy model to CPU after inference.""" futures = self.worker_group.run_all_workers_single_data("finish_inference") diff --git a/nemo_rl/models/policy/utils.py b/nemo_rl/models/policy/utils.py index a421d9bda6f..e7948493c51 100644 --- a/nemo_rl/models/policy/utils.py +++ b/nemo_rl/models/policy/utils.py @@ -359,7 +359,12 @@ def calculate_aligned_size(size_bytes: int, alignment: int = 512) -> int: def stream_weights_via_ipc_zmq_impl( - params_generator, buffer_size_bytes: int, zmq_socket, rank: int, worker_name: str + params_generator, + buffer_size_bytes: int, + zmq_socket, + rank: int, + worker_name: str, + buffer_cache: dict[str, Any] | None = None, ) -> None: """Shared implementation for streaming weights via IPC ZMQ with improved memory management. @@ -372,6 +377,11 @@ def stream_weights_via_ipc_zmq_impl( zmq_socket: ZMQ socket for communication rank: Worker rank for logging worker_name: Name of the worker for logging + buffer_cache: Optional caller-owned dict that keeps the ping-pong staging + buffers alive across calls. When provided, buffers are reused between + refits (skipping two large allocations plus a gc.collect/empty_cache + pair per refit) as long as the buffer size and device are unchanged. + The buffers then stay resident on the GPU between refits. """ # Divide total buffer size by 2 because we use two individual buffers (ping-pong) for overlapping communication. buffer_size_bytes = buffer_size_bytes // 2 @@ -419,6 +429,11 @@ def release_staging_buffers() -> None: """Release acyclic IPC buffers without scanning the worker object graph.""" nonlocal buffer_a, buffer_b, current_buffer + if buffer_cache is not None: + # Persistent-buffer mode: the buffers intentionally outlive this + # call (and the refit), so both the mid-stream reclaim and the + # final cleanup are skipped. + return had_buffers = buffer_a is not None or buffer_b is not None current_buffer = None buffer_a = None @@ -438,8 +453,30 @@ def release_staging_buffers() -> None: buffer_device = tensor.device if buffer_device.type == "cpu" and torch.cuda.is_available(): buffer_device = torch.device("cuda", torch.cuda.current_device()) - buffer_a = allocate_buffer(buffer_device) - buffer_b = allocate_buffer(buffer_device) + if ( + buffer_cache is not None + and "a" in buffer_cache + and buffer_cache.get("device") == buffer_device + ): + # Reuse the first-refit size even when dynamic sizing asks + # for more: parameters are identical every refit, and with + # vLLM sleep level 2 the free-memory-based request inflates + # while rollout weights are discarded - allocating (and + # keeping) larger buffers then OOMs vLLM's wake_up remap. + buffer_a = buffer_cache["a"] + buffer_b = buffer_cache["b"] + buffer_size_bytes = buffer_cache["size"] + else: + buffer_a = allocate_buffer(buffer_device) + buffer_b = allocate_buffer(buffer_device) + if buffer_cache is not None: + buffer_cache.clear() + buffer_cache.update( + a=buffer_a, + b=buffer_b, + size=buffer_size_bytes, + device=buffer_device, + ) current_buffer = buffer_a aligned_size = calculate_aligned_size(tensor.nbytes) diff --git a/nemo_rl/models/policy/workers/dtensor_policy_worker.py b/nemo_rl/models/policy/workers/dtensor_policy_worker.py index d132d1c111f..1803c371852 100644 --- a/nemo_rl/models/policy/workers/dtensor_policy_worker.py +++ b/nemo_rl/models/policy/workers/dtensor_policy_worker.py @@ -257,6 +257,9 @@ def __init__( configure_dynamo_cache() self.cfg = config + # Staging-buffer cache for refit weight streaming; only populated when + # cfg["refit_persistent_ipc_buffers"] is enabled. + self._refit_ipc_buffer_cache: dict[str, Any] = {} # torch distributed init. Envars for rank, world_size, and master_addr and master_port are set from the ray remote call torch.distributed.init_process_group(backend="nccl") self.rank = torch.distributed.get_rank() @@ -1885,6 +1888,11 @@ def stream_weights_via_ipc_zmq( zmq_socket=self.zmq_socket, rank=self.rank, worker_name=str(self), + buffer_cache=( + self._refit_ipc_buffer_cache + if self.cfg.get("refit_persistent_ipc_buffers") + else None + ), ) def _checkpoint_engine_params( diff --git a/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py b/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py index 8ed9d584fed..93856922579 100644 --- a/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py +++ b/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py @@ -252,6 +252,9 @@ def __init__( # Store configuration self.cfg = config + # Staging-buffer cache for refit weight streaming; only populated when + # cfg["refit_persistent_ipc_buffers"] is enabled. + self._refit_ipc_buffer_cache: dict[str, Any] = {} # Reconstruct tokenizer/processor locally to avoid pickling across # incompatible transformers versions (v4 head node → v5 worker). @@ -1130,6 +1133,11 @@ def stream_weights_via_ipc_zmq( zmq_socket=self.zmq_socket, rank=self.rank, worker_name=str(self), + buffer_cache=( + self._refit_ipc_buffer_cache + if self.cfg.get("refit_persistent_ipc_buffers") + else None + ), ) @torch.no_grad() diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index eecd310a782..666ce2ceb37 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -301,6 +301,19 @@ def __init__( self.cfg = config self._router_replay_enabled = router_replay_enabled(config) + # Staging-buffer cache for refit weight streaming; only populated when + # cfg["refit_persistent_ipc_buffers"] is enabled. + self._refit_ipc_buffer_cache: dict[str, Any] = {} + # HF param names to MXFP8-quantize on the trainer during refit; set via + # enable_refit_prequantize() when vllm_cfg.refit_prequantize is on. + self._refit_prequant_names: set[str] = set() + # Pinned host staging for the reference-policy swap; only populated when + # megatron_cfg["pinned_reference_swap"] is enabled. Buffer contents are + # only live within a single use_reference_model call (every copy + # synchronizes before control leaves it) and each call re-reads + # model.state_dict(), so offload_before/after_refit moving or replacing + # param storages between calls cannot collide with these buffers. + self._pinned_swap_save_buffers: dict[str, torch.Tensor] = {} self._nixl_preinit_agent = maybe_preinit_nixl_checkpoint_engine(config) # Set rank for non-collocated to check which ranks to broadcast from @@ -1531,6 +1544,7 @@ def _apply_state_dict_to_model( source_state_dict: dict, *, raise_if_key_missing: bool = False, + non_blocking: bool = False, ) -> None: """Apply a state dict to self.model in-place. @@ -1542,6 +1556,8 @@ def _apply_state_dict_to_model( source_state_dict: State dict to apply (e.g. reference_state_dict or saved model_state_dict). raise_if_key_missing: If True, raise when a key in self.model.state_dict() is missing from source_state_dict; if False, skip such keys. + non_blocking: Passed to the in-place copies; callers staging from + pinned CPU memory must synchronize afterwards. """ for state_dict_key, param_or_buf in self.model.state_dict().items(): if ( @@ -1562,7 +1578,7 @@ def _apply_state_dict_to_model( isinstance(source_value, torch.Tensor) and param_or_buf.shape == source_value.shape ): - param_or_buf.copy_(source_value) + param_or_buf.copy_(source_value, non_blocking=non_blocking) continue # Case 2: _extra_state (shape mismatch or non-Tensor) → set_extra_state() @@ -1593,12 +1609,39 @@ def use_reference_model(self): self.disable_forward_pre_hook() with torch.no_grad(): + # NotRequired key: absent means disabled, default lives in the exemplar YAML. + use_pinned_swap = bool( + self.cfg["megatron_cfg"].get("pinned_reference_swap") + ) + # Save original references model_state_dict = {} for name, item in self.model.state_dict().items(): if isinstance(item, torch.Tensor): - item = item.detach().to(device="cpu", non_blocking=True, copy=True) + # extra_state tensors stay on fresh pageable copies: + # set_extra_state() may retain the tensor it is given, and + # a reused pinned buffer would mutate it on the next swap. + if use_pinned_swap and "extra_state" not in name: + buf = self._pinned_swap_save_buffers.get(name) + if buf is None: + buf = torch.empty( + item.shape, + dtype=item.dtype, + device="cpu", + pin_memory=True, + ) + self._pinned_swap_save_buffers[name] = buf + buf.copy_(item.detach(), non_blocking=True) + item = buf + else: + item = item.detach().to( + device="cpu", non_blocking=True, copy=True + ) model_state_dict[name] = item + if use_pinned_swap: + # D2H saves must land before the reference apply overwrites + # the params they read from. + torch.cuda.synchronize() # Swap reference state into self.model. Use _apply_state_dict_to_model # (rather than load_state_dict) so FP8 _extra_state with mismatched shape @@ -1606,7 +1649,10 @@ def use_reference_model(self): self._apply_state_dict_to_model( self.reference_state_dict, raise_if_key_missing=True, + non_blocking=use_pinned_swap, ) + if use_pinned_swap: + torch.cuda.synchronize() if self.cfg["megatron_cfg"]["empty_unused_memory_level"] >= 1: gc.collect() @@ -1638,7 +1684,10 @@ def use_reference_model(self): self._apply_state_dict_to_model( model_state_dict, raise_if_key_missing=True, + non_blocking=use_pinned_swap, ) + if use_pinned_swap: + torch.cuda.synchronize() if self.cfg["megatron_cfg"]["empty_unused_memory_level"] >= 1: gc.collect() @@ -1753,6 +1802,44 @@ def prepare_refit_info(self) -> None: return refit_param_info_hf + def enable_refit_prequantize(self, param_names: list[str]) -> dict[str, Any]: + """Quantize the listed HF params to MXFP8 on the trainer during refit. + + Args: + param_names: fp8-eligible parameter names reported by the vLLM + workers (see VllmInternalWorkerExtension.prepare_refit_info). + + Returns: + Updated refit metadata: the listed params become float8_e4m3fn and + each gains a *_scale_from_checkpoint uint8 entry. + """ + self._refit_prequant_names = set(param_names) + + refit_param_info_hf = {} + for name, tensor in self._iter_params_with_optional_kv_scales(): + refit_param_info_hf[name] = (tensor.shape, tensor.dtype) + return refit_param_info_hf + + def _maybe_prequantize_param( + self, name: str, tensor: torch.Tensor + ) -> Iterator[tuple[str, torch.Tensor]]: + if ( + name not in self._refit_prequant_names + or tensor.dtype == torch.float8_e4m3fn + ): + yield name, tensor + return + + # Deferred: pulls in the heavy nemo_rl...generation.vllm package init, + # which trainer workers only need when prequantized refit is enabled. + from nemo_rl.models.generation.vllm.quantization.fp8_train_utils import ( + mxfp8_e4m3_quantize_for_refit, + ) + + param_lp, param_scale = mxfp8_e4m3_quantize_for_refit(tensor) + yield name, param_lp + yield name + "_scale_from_checkpoint", param_scale + def _collect_mtp_metrics( self, metrics: dict[str, Any], @@ -1936,9 +2023,10 @@ def _iter_params_with_optional_kv_scales( conversion_tasks=self.refit_conversion_tasks, ) - # Yield the original parameters first. + # Yield the original parameters first, MXFP8-quantizing on the trainer + # when pre-quantized refit is enabled for the parameter. for name, tensor in base_iter: - yield name, tensor + yield from self._maybe_prequantize_param(name, tensor) if self.draft_model is not None: from nemo_rl.models.megatron.draft import export_eagle_weights_to_hf @@ -2002,6 +2090,11 @@ def stream_weights_via_ipc_zmq( zmq_socket=self.zmq_socket, rank=self.rank, worker_name=str(self), + buffer_cache=( + self._refit_ipc_buffer_cache + if self.cfg.get("refit_persistent_ipc_buffers") + else None + ), ) @torch.no_grad() @@ -2113,56 +2206,7 @@ def offload_before_refit(self): self._clear_fp8_caches() if self.cfg["megatron_cfg"].get("clear_memory_caches_before_refit", False): - # Clear RotaryEmbedding's @lru_cache(maxsize=32). The cache accumulates one - # entry per unique (max_seq_len, offset, packed_seq) seen, and each entry is - # a GPU tensor (the concatenated sin/cos embedding). With training + logprob - # runs at different sequence lengths, the cache fills quickly and the tensors - # anchor large CUDA segments. - try: - from megatron.core.models.common.embeddings.rotary_pos_embedding import ( - RotaryEmbedding, - ) - - RotaryEmbedding.forward.cache_clear() - except Exception: - pass - - # Clear MoE token dispatcher persistent routing tensors. - # - # MoETokenDispatcher is a plain Python class (NOT an nn.Module), so iterating - # self.model.modules() never yields it. We must access it via the token_dispatcher - # attribute on MoELayer nn.Module objects. - # - # When recompute_mlp=True and fp8=True, - # transformer_layer._forward_mlp wraps self.mlp (the MoE layer) with te_checkpoint. - # te_checkpoint._CheckpointFunction.backward recomputes the forward with - # torch.enable_grad(), which causes dispatch_preprocess to store - # dispatcher.probs = routing_probs (with grad_fn, under enable_grad) - # This creates a reference cycle: - # _CheckpointFunctionBackward → ctx → ctx.run_function=mlp - # → mlp.token_dispatcher.probs → probs.grad_fn → ... → _CheckpointFunctionBackward - # - # Breaking this cycle by nulling dispatcher.probs frees BOTH: - # - the routing tensors - # - the te_checkpoint ctx saved tensors - try: - for module in self.model.modules(): - if not hasattr(module, "token_dispatcher"): - continue - dispatcher = module.token_dispatcher - if dispatcher is None: - continue - for attr in ( - "probs", # AllToAll + AllGather - "routing_map", # AllToAll - "reversed_local_input_permutation_mapping", # AllToAll - "local_probs", # AllGather - "local_map", # AllGather - ): - if isinstance(getattr(dispatcher, attr, None), torch.Tensor): - setattr(dispatcher, attr, None) - except Exception: - pass + self._clear_rope_and_moe_dispatcher_caches() torch.randn(1).cuda() # wake up torch allocator if ( @@ -2183,6 +2227,59 @@ def offload_before_refit(self): ) no_grad.__exit__(None, None, None) + def _clear_rope_and_moe_dispatcher_caches(self) -> None: + """Clear rotary-embedding and MoE dispatcher caches repopulated by forwards.""" + # Clear RotaryEmbedding's @lru_cache(maxsize=32). The cache accumulates one + # entry per unique (max_seq_len, offset, packed_seq) seen, and each entry is + # a GPU tensor (the concatenated sin/cos embedding). With training + logprob + # runs at different sequence lengths, the cache fills quickly and the tensors + # anchor large CUDA segments. + try: + from megatron.core.models.common.embeddings.rotary_pos_embedding import ( + RotaryEmbedding, + ) + + RotaryEmbedding.forward.cache_clear() + except Exception: + pass + + # Clear MoE token dispatcher persistent routing tensors. + # + # MoETokenDispatcher is a plain Python class (NOT an nn.Module), so iterating + # self.model.modules() never yields it. We must access it via the token_dispatcher + # attribute on MoELayer nn.Module objects. + # + # When recompute_mlp=True and fp8=True, + # transformer_layer._forward_mlp wraps self.mlp (the MoE layer) with te_checkpoint. + # te_checkpoint._CheckpointFunction.backward recomputes the forward with + # torch.enable_grad(), which causes dispatch_preprocess to store + # dispatcher.probs = routing_probs (with grad_fn, under enable_grad) + # This creates a reference cycle: + # _CheckpointFunctionBackward → ctx → ctx.run_function=mlp + # → mlp.token_dispatcher.probs → probs.grad_fn → ... → _CheckpointFunctionBackward + # + # Breaking this cycle by nulling dispatcher.probs frees BOTH: + # - the routing tensors + # - the te_checkpoint ctx saved tensors + try: + for module in self.model.modules(): + if not hasattr(module, "token_dispatcher"): + continue + dispatcher = module.token_dispatcher + if dispatcher is None: + continue + for attr in ( + "probs", # AllToAll + AllGather + "routing_map", # AllToAll + "reversed_local_input_permutation_mapping", # AllToAll + "local_probs", # AllGather + "local_map", # AllGather + ): + if isinstance(getattr(dispatcher, attr, None), torch.Tensor): + setattr(dispatcher, attr, None) + except Exception: + pass + @wrap_with_nvtx_name("megatron_policy_worker/offload_after_refit") def offload_after_refit(self): """Offload as much as possible on the CPU.""" @@ -2191,7 +2288,27 @@ def offload_after_refit(self): self.model = self.move_model(self.model, "cpu") self.model.eval() torch.randn(1).cuda() # wake up torch allocator - self.offload_before_refit() # rerun the old offload function + if self.cfg["megatron_cfg"].get("refit_slim_offload_after"): + # Grad buffers were already offloaded by offload_before_refit at + # the start of the refit, so skip the full rerun (grad-buffer moves + # and a second gc/empty_cache pair). Cache clears must still honor + # their knobs: callers may have run a forward pass since (e.g. + # teacher logits in distillation), repopulating TE fp8 workspaces, + # the rotary-embedding lru_cache, and MoE dispatcher tensors. + if self.fp8_cfg and self.fp8_cfg.get("force_clear_fp8_caches", False): + self._clear_fp8_caches() + if self.cfg["megatron_cfg"].get("clear_memory_caches_before_refit", False): + self._clear_rope_and_moe_dispatcher_caches() + if ( + hasattr(self, "optimizer") + and self.optimizer is not None + and not self.optimizer_cpu_offload + ): + self.move_optimizer("cpu") + gc.collect() + torch.cuda.empty_cache() + else: + self.offload_before_refit() # rerun the old offload function allocated = torch.cuda.memory_allocated() / (1024**3) # Convert to GB reserved = torch.cuda.memory_reserved() / (1024**3) # Convert to GB diff --git a/pyrefly.toml b/pyrefly.toml index 7f7a6ccccb0..52d20652f70 100644 --- a/pyrefly.toml +++ b/pyrefly.toml @@ -7,6 +7,7 @@ replace-imports-with-any = [ "datasets.*", "transformers.*", "vllm.*", + "flashinfer.*", "math_verify.*", "sympy.*", "torchdata.*", diff --git a/tests/unit/models/generation/test_mxfp8_prequant.py b/tests/unit/models/generation/test_mxfp8_prequant.py new file mode 100644 index 00000000000..f6f0af25255 --- /dev/null +++ b/tests/unit/models/generation/test_mxfp8_prequant.py @@ -0,0 +1,149 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import pytest +import torch + +from nemo_rl.models.generation.vllm.quantization.fp8_train_utils import ( + MXFP8_BLOCK_SIZE, + _mxfp8_e4m3_quantize_torch, + mxfp8_e4m3_quantize_for_refit, +) + + +def _dequantize(x_fp8: torch.Tensor, scales: torch.Tensor) -> torch.Tensor: + num_blocks = x_fp8.shape[-1] // MXFP8_BLOCK_SIZE + x_blocked = x_fp8.to(torch.float32).view( + *x_fp8.shape[:-1], num_blocks, MXFP8_BLOCK_SIZE + ) + descale = torch.exp2(scales.to(torch.float32) - 127.0) + return (x_blocked * descale.unsqueeze(-1)).view(*x_fp8.shape) + + +@pytest.mark.parametrize("shape", [(64, 128), (7, 96), (4, 16, 64)]) +def test_torch_reference_shapes_and_roundtrip(shape): + torch.manual_seed(0) + x = torch.randn(*shape, dtype=torch.bfloat16) + + x_fp8, scales = _mxfp8_e4m3_quantize_torch(x) + + assert x_fp8.shape == x.shape + assert x_fp8.dtype == torch.float8_e4m3fn + assert scales.dtype == torch.uint8 + expected_scale_shape = (*shape[:-1], shape[-1] // MXFP8_BLOCK_SIZE) + assert tuple(scales.shape) == expected_scale_shape + + x_dq = _dequantize(x_fp8, scales) + x32 = x.to(torch.float32) + abs_err = (x_dq - x32).abs() + block_amax = ( + x32.abs() + .reshape(*shape[:-1], shape[-1] // MXFP8_BLOCK_SIZE, MXFP8_BLOCK_SIZE) + .amax(dim=-1, keepdim=True) + .expand(*shape[:-1], shape[-1] // MXFP8_BLOCK_SIZE, MXFP8_BLOCK_SIZE) + .reshape(shape) + ) + # e4m3 has 3 mantissa bits: elements within the block's representable range + # must round-trip to ~12.5% relative error; elements far below the block + # amax may legitimately quantize to zero, so bound them by absolute error + # instead of a ratio (keeps the test deterministic across torch versions). + representable = x32.abs() >= block_amax / 64 + rel_err = (abs_err / x32.abs().clamp(min=1e-6))[representable] + assert rel_err.median() < 0.05 + assert rel_err.max() < 0.25 + assert (abs_err[~representable] <= block_amax[~representable] / 32).all() + + +def test_last_dim_not_divisible_raises(): + x = torch.randn(8, MXFP8_BLOCK_SIZE + 1, dtype=torch.bfloat16) + with pytest.raises(AssertionError): + _mxfp8_e4m3_quantize_torch(x) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_refit_quantize_matches_receiver_path(): + """Bitwise parity with the vLLM receiver path (mxfp8_e4m3_quantize + squeeze).""" + vllm_mxfp8 = pytest.importorskip( + "vllm.model_executor.layers.quantization.utils.mxfp8_utils" + ) + + torch.manual_seed(0) + x = torch.randn(256, 512, dtype=torch.bfloat16, device="cuda") + + ref_lp, ref_scale = vllm_mxfp8.mxfp8_e4m3_quantize(x) + ref_scale = torch.squeeze(ref_scale, dim=-1) + + got_lp, got_scale = mxfp8_e4m3_quantize_for_refit(x) + + assert got_lp.dtype == ref_lp.dtype + assert torch.equal(got_lp.view(torch.uint8), ref_lp.view(torch.uint8)) + assert got_scale.dtype == ref_scale.dtype + assert got_scale.shape == ref_scale.shape + assert torch.equal(got_scale.reshape(-1), ref_scale.reshape(-1)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize( + "is_gated,intermediate_size,hidden_size", + [ + # Aligned: both scale K dims (hidden/32=8, intermediate/32=4) are %4. + (True, 128, 256), + # w2 scale K = 192/32 = 6, so pad_flashinfer_scale_k pads it to 8. + (True, 192, 128), + # Non-gated (single w13 shard), aligned. + (False, 128, 256), + ], +) +def test_batched_moe_shuffle_matches_per_expert( + is_gated, intermediate_size, hidden_size +): + """Bitwise parity of the batched TRTLLM MoE shuffle with the per-expert loop.""" + pytest.importorskip("flashinfer") + fp8 = pytest.importorskip("nemo_rl.models.generation.vllm.quantization.fp8") + + from types import SimpleNamespace + + torch.manual_seed(0) + num_experts = 4 + w13_rows = (2 if is_gated else 1) * intermediate_size + + def rand_bytes(*shape): + return torch.randint(0, 256, shape, dtype=torch.uint8, device="cuda") + + w13_weight = rand_bytes(num_experts, w13_rows, hidden_size).view( + torch.float8_e4m3fn + ) + w2_weight = rand_bytes(num_experts, hidden_size, intermediate_size).view( + torch.float8_e4m3fn + ) + w13_scale = rand_bytes(num_experts, w13_rows, hidden_size // MXFP8_BLOCK_SIZE) + w2_scale = rand_bytes( + num_experts, hidden_size, intermediate_size // MXFP8_BLOCK_SIZE + ) + + layer = SimpleNamespace() # holds the cached row permutations + epilogue_tile_m = 128 + batched = fp8._shuffle_mxfp8_moe_batched( + layer, w13_weight, w2_weight, w13_scale, w2_scale, is_gated, epilogue_tile_m + ) + reference = fp8._shuffle_mxfp8_moe_per_expert( + w13_weight, w2_weight, w13_scale, w2_scale, is_gated, epilogue_tile_m + ) + + for got, want, name in zip( + batched, reference, ("w13_weight", "w2_weight", "w13_scale", "w2_scale") + ): + assert got.shape == want.shape, name + assert got.dtype == want.dtype, name + assert torch.equal(got.view(torch.uint8), want.view(torch.uint8)), name diff --git a/tests/unit/models/generation/test_vllm_backend.py b/tests/unit/models/generation/test_vllm_backend.py index 3144efb929b..4e7e0d447b2 100644 --- a/tests/unit/models/generation/test_vllm_backend.py +++ b/tests/unit/models/generation/test_vllm_backend.py @@ -127,7 +127,6 @@ def packed_broadcast_consumer(iterator, group, src, post_unpack_func): post_unpack_func([("model.weight", "weight-value")]) ext._load_weights = load_weights - ext._maybe_process_fp8_kv_cache = lambda: call_order.append("kv") monkeypatch.setattr( vllm_backend, "packed_broadcast_consumer", packed_broadcast_consumer ) @@ -147,7 +146,6 @@ def packed_broadcast_consumer(iterator, group, src, post_unpack_func): "config_enter", "process", "config_exit", - "kv", "gc", "empty_cache", ] diff --git a/tests/unit/models/generation/test_vllm_checkpoint_engine.py b/tests/unit/models/generation/test_vllm_checkpoint_engine.py index 75051899ff1..15f65356e4c 100644 --- a/tests/unit/models/generation/test_vllm_checkpoint_engine.py +++ b/tests/unit/models/generation/test_vllm_checkpoint_engine.py @@ -15,6 +15,8 @@ """Tests for vLLM checkpoint-engine worker lifecycle helpers.""" import asyncio +from collections.abc import Callable, Iterator +from contextlib import contextmanager from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock @@ -95,7 +97,20 @@ async def receive_weight_batches(self): worker._load_weights = lambda batch: events.append( ("load", [name for name, _weight in batch]) ) - worker._maybe_process_fp8_kv_cache = lambda: events.append(("fp8",)) + + @contextmanager + def weight_update_lifecycle( + transport: str, + ) -> Iterator[Callable[[], None]]: + events.append(("setup", transport)) + + def finalize() -> None: + events.append(("finalize",)) + + yield finalize + events.append(("teardown", transport)) + + worker._weight_update_lifecycle = weight_update_lifecycle monkeypatch.setattr( torch.cuda, "current_stream", @@ -104,11 +119,13 @@ async def receive_weight_batches(self): assert asyncio.run(worker._update_weights_from_checkpoint_engine_async()) is True assert events == [ + ("setup", "checkpoint_engine"), ("load", ["a"]), ("sync",), ("load", ["b", "c"]), ("sync",), - ("fp8",), + ("finalize",), + ("teardown", "checkpoint_engine"), ] diff --git a/tests/unit/models/generation/test_vllm_modelopt_real_quant_config.py b/tests/unit/models/generation/test_vllm_modelopt_real_quant_config.py index 4e71431790f..125f50b9ee3 100644 --- a/tests/unit/models/generation/test_vllm_modelopt_real_quant_config.py +++ b/tests/unit/models/generation/test_vllm_modelopt_real_quant_config.py @@ -1129,9 +1129,9 @@ def make_model(expert_map): batched_forwarded = [] extension = _make_real_quant_extension(backend, make_model(None), []) + _patch_real_quant_load(monkeypatch, backend, batched_forwarded) extension.prepare_refit_info(state_dict_info) extension._nrl_w13_num_shards_by_prefix = {prefix: 1} - _patch_real_quant_load(monkeypatch, backend, batched_forwarded) assert ( extension._load_weights( [ diff --git a/tests/unit/reference_configs/grpo_math_1B.yaml b/tests/unit/reference_configs/grpo_math_1B.yaml index 990230f9274..c29b48b3494 100644 --- a/tests/unit/reference_configs/grpo_math_1B.yaml +++ b/tests/unit/reference_configs/grpo_math_1B.yaml @@ -157,6 +157,8 @@ policy: megatron_cfg: enabled: false + refit_slim_offload_after: false + pinned_reference_swap: false checkpoint: async_save: true ckpt_assume_constant_structure: true @@ -300,6 +302,7 @@ policy: # makes the training sequence length divisible by the tensor parallel size # this is useful for sequence parallel training make_sequence_length_divisible_by: ${policy.dtensor_cfg.tensor_parallel_size} + refit_persistent_ipc_buffers: false max_grad_norm: 1.0 optimizer: From 0fb59b6f3cd7806bc6ec263dcb9f63e6340f123d Mon Sep 17 00:00:00 2001 From: sna Date: Mon, 27 Jul 2026 12:01:53 -0700 Subject: [PATCH 02/76] fix(trtllm): align refit metadata interface Signed-off-by: sna --- .../generation/trtllm/trtllm_generation.py | 7 +++++- .../trtllm/test_trtllm_generation.py | 24 +++++++++++++++++++ 2 files changed, 30 insertions(+), 1 deletion(-) diff --git a/nemo_rl/models/generation/trtllm/trtllm_generation.py b/nemo_rl/models/generation/trtllm/trtllm_generation.py index 978303b5b06..89f153b48a8 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_generation.py +++ b/nemo_rl/models/generation/trtllm/trtllm_generation.py @@ -435,13 +435,18 @@ def finish_generation(self, *args: Any, **kwargs: Any) -> bool: print(f"Error in finish_generation: {e}") return False - def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: + def prepare_refit_info( + self, state_dict_info: Optional[dict[str, Any]] + ) -> Optional[list[str]]: + if state_dict_info is None: + return None futures = self.worker_group.run_all_workers_single_data( "prepare_refit_info_async", state_dict_info=state_dict_info, run_rank_0_only_axes=["tensor_parallel"], ) ray.get(futures) + return None def start_gpu_profiling(self) -> None: """Grpo profiling protocol: start nsys capture on the GPU workers.""" diff --git a/tests/unit/models/generation/trtllm/test_trtllm_generation.py b/tests/unit/models/generation/trtllm/test_trtllm_generation.py index d9dc1eaea70..bd02c1c30fb 100644 --- a/tests/unit/models/generation/trtllm/test_trtllm_generation.py +++ b/tests/unit/models/generation/trtllm/test_trtllm_generation.py @@ -294,6 +294,30 @@ def test_generation_lifecycle_routes_by_colocation( ) +def test_prepare_refit_info_skips_missing_metadata_and_dispatches_present_metadata( + monkeypatch, +): + generation = _bare_generation() + ray_get = MagicMock() + monkeypatch.setattr(trtllm_generation.ray, "get", ray_get) + + assert generation.prepare_refit_info(None) is None + generation.worker_group.run_all_workers_single_data.assert_not_called() + ray_get.assert_not_called() + + state_dict_info = {"weight": {"shape": [2, 2]}} + futures = [SimpleNamespace()] + generation.worker_group.run_all_workers_single_data.return_value = futures + + assert generation.prepare_refit_info(state_dict_info) is None + generation.worker_group.run_all_workers_single_data.assert_called_once_with( + "prepare_refit_info_async", + state_dict_info=state_dict_info, + run_rank_0_only_axes=["tensor_parallel"], + ) + ray_get.assert_called_once_with(futures) + + @pytest.mark.parametrize( ("in_flight", "recompute_kv", "expected_drain"), [(False, False, True), (True, False, False), (True, True, False)], From c4a5e14a195bdecd7d6ebb8fcb479fd0d4aed97f Mon Sep 17 00:00:00 2001 From: sna Date: Mon, 27 Jul 2026 12:18:49 -0700 Subject: [PATCH 03/76] test(vllm): include MXFP8 refit tests in L0 Signed-off-by: sna --- tests/unit/models/generation/test_mxfp8_prequant.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/unit/models/generation/test_mxfp8_prequant.py b/tests/unit/models/generation/test_mxfp8_prequant.py index f6f0af25255..b8a83a4dafb 100644 --- a/tests/unit/models/generation/test_mxfp8_prequant.py +++ b/tests/unit/models/generation/test_mxfp8_prequant.py @@ -21,6 +21,8 @@ mxfp8_e4m3_quantize_for_refit, ) +pytestmark = pytest.mark.vllm + def _dequantize(x_fp8: torch.Tensor, scales: torch.Tensor) -> torch.Tensor: num_blocks = x_fp8.shape[-1] // MXFP8_BLOCK_SIZE From 5c53bd493af7eb7c09d02ed5c8dab82eabdb3f72 Mon Sep 17 00:00:00 2001 From: sna Date: Mon, 27 Jul 2026 12:32:54 -0700 Subject: [PATCH 04/76] fix(refit): complete prequant metadata handshake Signed-off-by: sna --- .../generation/sglang/sglang_generation.py | 6 +- .../models/generation/vllm/vllm_generation.py | 3 + .../checkpoint_engine_weight_synchronizer.py | 10 +++- .../sglang/test_sglang_generation.py | 6 ++ .../models/generation/test_vllm_generation.py | 9 +++ .../models/policy/test_megatron_worker.py | 60 ++++++++++++++++++- ...t_checkpoint_engine_weight_synchronizer.py | 21 ++++++- 7 files changed, 110 insertions(+), 5 deletions(-) diff --git a/nemo_rl/models/generation/sglang/sglang_generation.py b/nemo_rl/models/generation/sglang/sglang_generation.py index 9aa21547a63..0aff8acee37 100644 --- a/nemo_rl/models/generation/sglang/sglang_generation.py +++ b/nemo_rl/models/generation/sglang/sglang_generation.py @@ -731,8 +731,10 @@ def init_collective( ) -> list[ray.ObjectRef]: return [] - def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: - pass + def prepare_refit_info( + self, state_dict_info: dict[str, Any] | None + ) -> list[str] | None: + return None def update_weights_via_ipc_zmq(self) -> list[ray.ObjectRef]: return [] diff --git a/nemo_rl/models/generation/vllm/vllm_generation.py b/nemo_rl/models/generation/vllm/vllm_generation.py index 16ae57292ac..4708ef26581 100644 --- a/nemo_rl/models/generation/vllm/vllm_generation.py +++ b/nemo_rl/models/generation/vllm/vllm_generation.py @@ -926,6 +926,9 @@ def prepare_refit_info( (vllm_cfg.refit_prequantize), the parameter names the engine wants quantized on the trainer before streaming. None otherwise. """ + if state_dict_info is None: + return None + # Choose the appropriate method based on async_engine setting method_name = ( "prepare_refit_info_async" diff --git a/nemo_rl/weight_sync/checkpoint_engine_weight_synchronizer.py b/nemo_rl/weight_sync/checkpoint_engine_weight_synchronizer.py index 9e035fb4ff3..76af722cb62 100644 --- a/nemo_rl/weight_sync/checkpoint_engine_weight_synchronizer.py +++ b/nemo_rl/weight_sync/checkpoint_engine_weight_synchronizer.py @@ -67,7 +67,15 @@ class CheckpointEngineWeightSynchronizer(WeightSynchronizer): _bucket_size_bytes: int | None = None def init_communicator(self) -> None: - self._generation.prepare_refit_info(self._policy.prepare_refit_info()) + state_dict_info = self._policy.prepare_refit_info() + prequant_names = self._generation.prepare_refit_info(state_dict_info) + if prequant_names: + updated_info = self._policy.enable_refit_prequantize(prequant_names) + if updated_info is None: + raise RuntimeError( + "Trainer-side refit prequantization did not return updated metadata." + ) + self._generation.prepare_refit_info(updated_info) self._ensure_checkpoint_engine_ready() @property diff --git a/tests/unit/models/generation/sglang/test_sglang_generation.py b/tests/unit/models/generation/sglang/test_sglang_generation.py index 2e69c418c99..a365fef8c38 100644 --- a/tests/unit/models/generation/sglang/test_sglang_generation.py +++ b/tests/unit/models/generation/sglang/test_sglang_generation.py @@ -46,6 +46,12 @@ pytestmark = pytest.mark.sglang +def test_prepare_refit_info_accepts_missing_metadata(): + generation = SGLangGeneration.__new__(SGLangGeneration) + + assert generation.prepare_refit_info(None) is None + + @pytest.fixture(scope="module") def ray_cluster(): """Initialise Ray once for this module's tests.""" diff --git a/tests/unit/models/generation/test_vllm_generation.py b/tests/unit/models/generation/test_vllm_generation.py index d402e682ee9..69f8b6ee358 100644 --- a/tests/unit/models/generation/test_vllm_generation.py +++ b/tests/unit/models/generation/test_vllm_generation.py @@ -138,6 +138,15 @@ } +@pytest.mark.vllm +def test_prepare_refit_info_skips_missing_metadata(): + generation = VllmGeneration.__new__(VllmGeneration) + generation.worker_group = MagicMock() + + assert generation.prepare_refit_info(None) is None + generation.worker_group.run_all_workers_single_data.assert_not_called() + + def test_resolve_enable_prefix_caching_respects_explicit_config(monkeypatch): def raise_if_called(): raise AssertionError("CUDA capability should not be queried") diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index 9727b50d06b..80c7df8364c 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -17,7 +17,7 @@ import time from pathlib import Path from types import SimpleNamespace -from typing import Optional +from typing import Any, Optional import numpy as np import pytest @@ -78,6 +78,64 @@ def test_megatron_move_model_does_not_serialize_extra_state(): assert model.scale.device.type == "cpu" +def test_checkpoint_engine_prequant_handshake_exports_mxfp8_weights(): + from nemo_rl.models.policy.workers.checkpoint_engine import ( + MegatronCheckpointEngineSendMixin, + ) + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + from nemo_rl.weight_sync.checkpoint_engine_weight_synchronizer import ( + CheckpointEngineWeightSynchronizer, + ) + + class _PrequantCheckpointWorker(MegatronCheckpointEngineSendMixin): + enable_refit_prequantize = MegatronPolicyWorkerImpl.enable_refit_prequantize + _iter_params_with_optional_kv_scales = ( + MegatronPolicyWorkerImpl._iter_params_with_optional_kv_scales + ) + _maybe_prequantize_param = MegatronPolicyWorkerImpl._maybe_prequantize_param + + name = "model.layers.0.mlp.down_proj.weight" + weight = torch.randn(64, 64, dtype=torch.bfloat16) + worker = _PrequantCheckpointWorker() + worker._refit_prequant_names = set() + worker.model = object() + worker.draft_model = None + worker.refit_conversion_tasks = [] + worker.cfg = {} + worker.megatron_bridge = SimpleNamespace( + export_hf_weights=lambda *_args, **_kwargs: iter([(name, weight)]) + ) + worker.prepare_refit_info = lambda: {name: (weight.shape, weight.dtype)} + worker.checkpoint_engine = SimpleNamespace(get_target_weight_layout=lambda: None) + + class _Generation: + def __init__(self) -> None: + self.refit_info: list[dict[str, Any]] = [] + + def prepare_refit_info( + self, state_dict_info: dict[str, Any] | None + ) -> list[str] | None: + assert state_dict_info is not None + self.refit_info.append(state_dict_info) + return [name] if len(self.refit_info) == 1 else None + + generation = _Generation() + synchronizer = CheckpointEngineWeightSynchronizer(worker, generation, {}) + synchronizer._ensure_checkpoint_engine_ready = lambda: None + + synchronizer.init_communicator() + exported = dict(worker._checkpoint_engine_weight_iterator()) + + scale_name = name + "_scale_from_checkpoint" + assert generation.refit_info[1][name][1] == torch.float8_e4m3fn + assert generation.refit_info[1][scale_name][1] == torch.uint8 + assert exported[name].dtype == torch.float8_e4m3fn + assert exported[scale_name].dtype == torch.uint8 + assert exported[scale_name].shape == (64, 2) + + def test_megatron_prepare_for_training_restores_optimizer(): from nemo_rl.models.policy.workers.megatron_policy_worker import ( MegatronPolicyWorkerImpl, diff --git a/tests/unit/weight_sync/test_checkpoint_engine_weight_synchronizer.py b/tests/unit/weight_sync/test_checkpoint_engine_weight_synchronizer.py index 8a02257e8fa..ecb8585acb9 100644 --- a/tests/unit/weight_sync/test_checkpoint_engine_weight_synchronizer.py +++ b/tests/unit/weight_sync/test_checkpoint_engine_weight_synchronizer.py @@ -14,7 +14,7 @@ """Tests for checkpoint-engine weight synchronization and factory routing.""" -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock, call, patch import pytest @@ -153,6 +153,25 @@ def _checkpoint_sync( class TestCheckpointEngineWeightSynchronizer: + def test_init_communicator_completes_prequant_handshake(self): + sync = _checkpoint_sync(MagicMock()) + sync._ensure_checkpoint_engine_ready = MagicMock() + state_dict_info = {"layer_0": {"shape": [4096, 4096]}} + prequant_names = ["layer_0"] + updated_info = {"layer_0": {"shape": [4096, 4096], "dtype": "float8_e4m3fn"}} + sync._policy.prepare_refit_info.return_value = state_dict_info + sync._generation.prepare_refit_info.side_effect = [prequant_names, None] + sync._policy.enable_refit_prequantize.return_value = updated_info + + sync.init_communicator() + + sync._policy.enable_refit_prequantize.assert_called_once_with(prequant_names) + assert sync._generation.prepare_refit_info.call_args_list == [ + call(state_dict_info), + call(updated_info), + ] + sync._ensure_checkpoint_engine_ready.assert_called_once_with() + @patch("nemo_rl.weight_sync.checkpoint_engine_weight_synchronizer.ray") def test_bucket_uses_minimum_total_memory_and_is_cached(self, mock_ray, capsys): config = _checkpoint_engine_cfg(bucket_memory_ratio=0.125) From 98fc19757c41560f476c52342108c8dab348d396 Mon Sep 17 00:00:00 2001 From: sna Date: Mon, 27 Jul 2026 15:45:51 -0700 Subject: [PATCH 05/76] test(refit): cover MXFP8 optimization paths Signed-off-by: sna --- tests/unit/algorithms/test_utils.py | 63 +++++ .../models/generation/test_vllm_backend.py | 98 ++++++++ .../generation/test_vllm_fp8_quantization.py | 62 +++++ .../generation/test_vllm_refit_loader.py | 151 ++++++++++++ .../models/policy/test_megatron_worker.py | 215 ++++++++++++++++++ .../models/policy/test_policy_validation.py | 24 ++ ...t_checkpoint_engine_weight_synchronizer.py | 17 ++ 7 files changed, 630 insertions(+) diff --git a/tests/unit/algorithms/test_utils.py b/tests/unit/algorithms/test_utils.py index d4523496488..a22decdf5e0 100755 --- a/tests/unit/algorithms/test_utils.py +++ b/tests/unit/algorithms/test_utils.py @@ -14,6 +14,7 @@ import math from datetime import datetime +from unittest.mock import MagicMock, call import pytest import torch @@ -24,6 +25,7 @@ WALL_CLOCK_EFFICIENCY_CATEGORIES, calculate_baseline_and_std_per_prompt, get_tokenizer, + maybe_enable_refit_prequantize, maybe_pad_last_batch, print_efficiency_summary, print_performance_metrics, @@ -737,3 +739,64 @@ def test_wall_waste_clamped_to_wall_time(self): assert result["efficiency/total_waste_s"] == 60.0 assert result["efficiency/productive_time_s"] == 0.0 assert result["efficiency/efficiency_pct"] == 0.0 + + +class TestMaybeEnableRefitPrequantize: + def test_returns_when_generation_does_not_request_prequantization(self): + policy = MagicMock() + generation = MagicMock() + generation.prepare_refit_info.return_value = None + state_dict_info = {"model.weight": ((2, 2), torch.bfloat16)} + + maybe_enable_refit_prequantize( + policy, + generation, + state_dict_info, + {"megatron_cfg": {"enabled": True}}, + ) + + generation.prepare_refit_info.assert_called_once_with(state_dict_info) + policy.enable_refit_prequantize.assert_not_called() + + @pytest.mark.parametrize( + "policy_config", + [{}, {"megatron_cfg": {"enabled": False}}], + ) + def test_rejects_prequantization_without_megatron(self, policy_config): + policy = MagicMock() + generation = MagicMock() + generation.prepare_refit_info.return_value = ["model.weight"] + + with pytest.raises(ValueError, match="requires the Megatron policy backend"): + maybe_enable_refit_prequantize( + policy, + generation, + {"model.weight": ((2, 2), torch.bfloat16)}, + policy_config, + ) + + policy.enable_refit_prequantize.assert_not_called() + + def test_refreshes_generation_metadata_after_prequantization(self): + policy = MagicMock() + generation = MagicMock() + state_dict_info = {"model.weight": ((2, 2), torch.bfloat16)} + updated_info = { + "model.weight": ((2, 2), torch.float8_e4m3fn), + "model.weight_scale_from_checkpoint": ((2, 1), torch.uint8), + } + generation.prepare_refit_info.side_effect = [["model.weight"], None] + policy.enable_refit_prequantize.return_value = updated_info + + maybe_enable_refit_prequantize( + policy, + generation, + state_dict_info, + {"megatron_cfg": {"enabled": True}}, + ) + + policy.enable_refit_prequantize.assert_called_once_with(["model.weight"]) + assert generation.prepare_refit_info.call_args_list == [ + call(state_dict_info), + call(updated_info), + ] diff --git a/tests/unit/models/generation/test_vllm_backend.py b/tests/unit/models/generation/test_vllm_backend.py index 80c44f65d97..4b892ef178d 100644 --- a/tests/unit/models/generation/test_vllm_backend.py +++ b/tests/unit/models/generation/test_vllm_backend.py @@ -119,6 +119,104 @@ def _make_mtp_refit_extension( return ext, drafter_model +@pytest.mark.vllm +@pytest.mark.parametrize("enabled", [False, True]) +def test_prepare_refit_info_reports_only_fp8_weights(monkeypatch, enabled): + from nemo_rl.models.generation.vllm import vllm_backend + from nemo_rl.models.generation.vllm.quantization import fp8 + + ext = vllm_backend.VllmInternalWorkerExtension.__new__( + vllm_backend.VllmInternalWorkerExtension + ) + model = object() + config = object() + ext.model_runner = SimpleNamespace(model=model, vllm_config=config) + state_dict_info = { + "model.linear.weight": ((2, 2), torch.bfloat16), + "model.norm.weight": ((2,), torch.bfloat16), + } + is_fp8_model = MagicMock(return_value=True) + checked_names = [] + + def is_fp8_weight(name, candidate_model): + assert candidate_model is model + checked_names.append(name) + return name == "model.linear.weight" + + monkeypatch.setattr( + fp8, + "global_fp8_config", + SimpleNamespace(is_mx=True, refit_prequantize=enabled), + ) + monkeypatch.setattr(fp8, "is_fp8_model", is_fp8_model) + monkeypatch.setattr(fp8, "_is_fp8_weight", is_fp8_weight) + + result = ext.prepare_refit_info(state_dict_info) + + assert ext.state_dict_info is state_dict_info + if enabled: + assert result == ["model.linear.weight"] + is_fp8_model.assert_called_once_with(config) + assert checked_names == list(state_dict_info) + else: + assert result is None + is_fp8_model.assert_not_called() + assert checked_names == [] + + +@pytest.mark.vllm +def test_sync_prepare_refit_info_unions_worker_names(): + from nemo_rl.models.generation.vllm.vllm_worker import ( + VllmGenerationWorkerImpl, + ) + + worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) + state_dict_info = {"model.weight": ((2, 2), torch.bfloat16)} + worker.llm = SimpleNamespace( + collective_rpc=MagicMock( + return_value=[ + None, + ["model.b.weight", "model.a.weight"], + ["model.a.weight"], + ] + ) + ) + + assert worker.prepare_refit_info(state_dict_info) == [ + "model.a.weight", + "model.b.weight", + ] + worker.llm.collective_rpc.assert_called_once_with( + "prepare_refit_info", + args=(state_dict_info,), + ) + + +@pytest.mark.vllm +@pytest.mark.asyncio +async def test_async_prepare_refit_info_unions_worker_names(): + from nemo_rl.models.generation.vllm.vllm_worker_async import ( + VllmAsyncGenerationWorkerImpl, + ) + + worker = VllmAsyncGenerationWorkerImpl.__new__(VllmAsyncGenerationWorkerImpl) + state_dict_info = {"model.weight": ((2, 2), torch.bfloat16)} + worker.llm = SimpleNamespace( + collective_rpc=AsyncMock( + return_value=[None, ["model.b.weight"], ["model.a.weight"]] + ) + ) + + assert await worker.prepare_refit_info_async(state_dict_info) == [ + "model.a.weight", + "model.b.weight", + ] + worker.llm.collective_rpc.assert_awaited_once_with( + "prepare_refit_info", + args=(state_dict_info,), + ) + + @pytest.mark.vllm @pytest.mark.parametrize("with_mtp", [False, True]) def test_update_weights_from_collective_processes_weights_after_loading( diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index ace5b0e8b91..cdc5589f97c 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -15,6 +15,7 @@ import types import pytest +import torch pytestmark = pytest.mark.vllm @@ -169,3 +170,64 @@ def fake_patch(path, _replacement): for path in patched_paths ) assert all(patcher.started for patcher in fp8.fp8_state.vllm_patches) + + +def test_load_weights_preserves_prequantized_mxfp8_and_clamps_scales( + fp8_module, monkeypatch +): + from nemo_rl.models.generation.vllm import vllm_backend + from vllm.model_executor.layers.quantization.utils import mxfp8_utils + + fp8 = fp8_module + fp8.global_fp8_config = types.SimpleNamespace(is_mx=True) + native = torch.ones(2, 2, dtype=torch.bfloat16) + prequantized = torch.ones(2, 2, dtype=torch.float8_e4m3fn) + receiver_quantized = torch.full((2, 2), 2.0, dtype=torch.bfloat16) + receiver_fp8 = torch.ones(2, 2, dtype=torch.float8_e4m3fn) + receiver_scales = torch.tensor([[[0], [7]], [[3], [0]]], dtype=torch.uint8) + loaded = [] + + monkeypatch.setattr( + fp8, + "_is_fp8_weight", + lambda name, _model: name != "model.native", + ) + monkeypatch.setattr( + mxfp8_utils, + "mxfp8_e4m3_quantize", + lambda tensor: ( + ( + receiver_fp8, + receiver_scales, + ) + if tensor is receiver_quantized + else pytest.fail("unexpected receiver quantization input") + ), + ) + monkeypatch.setattr( + vllm_backend, + "load_weights_maybe_cached", + lambda model, weights: loaded.extend(weights), + ) + model = object() + + fp8.load_weights( + [ + ("model.native", native), + ("model.prequantized.weight", prequantized), + ("model.receiver.weight", receiver_quantized), + ], + types.SimpleNamespace(model=model), + ) + + assert loaded[0][0] == "model.native" + assert loaded[0][1] is native + assert loaded[1][0] == "model.prequantized.weight" + assert loaded[1][1] is prequantized + assert loaded[2][0] == "model.receiver.weight" + assert loaded[2][1] is receiver_fp8 + assert loaded[3][0] == "model.receiver.weight_scale_from_checkpoint" + torch.testing.assert_close( + loaded[3][1], + torch.tensor([[1, 7], [3, 1]], dtype=torch.uint8), + ) diff --git a/tests/unit/models/generation/test_vllm_refit_loader.py b/tests/unit/models/generation/test_vllm_refit_loader.py index 999cf124e84..022d5c0aa91 100644 --- a/tests/unit/models/generation/test_vllm_refit_loader.py +++ b/tests/unit/models/generation/test_vllm_refit_loader.py @@ -138,6 +138,157 @@ def test_refit_load_weights_uses_full_weight_path_by_default(): assert loaded == [("model.weight", weight)] +@pytest.mark.vllm +def test_refit_loader_cache_records_replays_and_falls_back(monkeypatch): + from vllm.model_executor.model_loader.weight_utils import default_weight_loader + + from nemo_rl.models.generation.vllm.vllm_backend import ( + load_weights_maybe_cached, + ) + + events = [] + + def remote_loader(param, loaded_weight, *args, **kwargs): + events.append(("remote", loaded_weight, args, kwargs)) + return False + + def local_loader(param, loaded_weight, *args, **kwargs): + events.append(("local", loaded_weight, args, kwargs)) + with torch.no_grad(): + param.copy_(loaded_weight) + return None + + class Model: + def __init__(self): + self.remote = torch.nn.Parameter(torch.zeros(1), requires_grad=False) + self.local = torch.nn.Parameter(torch.zeros(1), requires_grad=False) + self.default = torch.nn.Parameter(torch.zeros(1), requires_grad=False) + self.remote.weight_loader = remote_loader + self.local.weight_loader = local_loader + self.default.weight_loader = default_weight_loader + self.load_calls = [] + + def named_parameters(self): + return [ + ("remote", self.remote), + ("local", self.local), + ("default", self.default), + ] + + def load_weights(self, *, weights): + self.load_calls.append([name for name, _weight in weights]) + loaded = set() + for name, weight in weights: + if name == "expert": + results = [ + self.remote.weight_loader( + self.remote, weight, "w1", expert_id=0 + ), + self.local.weight_loader(self.local, weight, "w1", expert_id=1), + ] + if any(result is not False for result in results): + loaded.add(name) + else: + self.default.weight_loader(self.default, weight) + loaded.add(name) + return loaded + + monkeypatch.setenv("NRL_REFIT_CACHED_LOADERS", "1") + model = Model() + first_expert = torch.tensor([1.0]) + first_default = torch.tensor([2.0]) + second_expert = torch.tensor([3.0]) + second_default = torch.tensor([4.0]) + + assert load_weights_maybe_cached( + model, + [("expert", first_expert), ("default", first_default)], + ) == {"expert", "default"} + assert load_weights_maybe_cached( + model, + [("expert", second_expert), ("default", second_default)], + ) == {"expert", "default"} + + cache = model._nrl_refit_loader_cache + assert model.load_calls == [["expert", "default"], ["default"]] + assert cache.uncached == {"default"} + assert set(cache.calls) == {"expert"} + assert len(cache.calls["expert"]) == 2 + first_loader, first_param, first_args, first_kwargs = cache.calls["expert"][0] + second_loader, second_param, second_args, second_kwargs = cache.calls["expert"][1] + assert first_loader is remote_loader + assert first_param is model.remote + assert first_args == ("w1",) + assert first_kwargs == {"expert_id": 0} + assert second_loader is local_loader + assert second_param is model.local + assert second_args == ("w1",) + assert second_kwargs == {"expert_id": 1} + assert [event[0] for event in events] == ["remote", "local", "remote", "local"] + assert events[2][1] is second_expert + assert events[3][1] is second_expert + assert model.remote.weight_loader is remote_loader + assert model.local.weight_loader is local_loader + torch.testing.assert_close(model.local, second_expert) + torch.testing.assert_close(model.default, second_default) + + +@pytest.mark.vllm +def test_refit_loader_cache_invalidates_replaced_parameter(monkeypatch): + from nemo_rl.models.generation.vllm.vllm_backend import ( + load_weights_maybe_cached, + ) + + events = [] + + def make_loader(label): + def loader(param, loaded_weight): + events.append(label) + with torch.no_grad(): + param.copy_(loaded_weight) + + return loader + + class Model: + def __init__(self): + self.remote = torch.nn.Parameter(torch.zeros(1), requires_grad=False) + self.local = torch.nn.Parameter(torch.zeros(1), requires_grad=False) + self.remote.weight_loader = make_loader("remote") + self.local.weight_loader = make_loader("local") + self.load_calls = [] + + def named_parameters(self): + return [("remote", self.remote), ("local", self.local)] + + def load_weights(self, *, weights): + self.load_calls.append([name for name, _weight in weights]) + for _name, weight in weights: + self.remote.weight_loader(self.remote, weight) + self.local.weight_loader(self.local, weight) + return {name for name, _weight in weights} + + monkeypatch.setenv("NRL_REFIT_CACHED_LOADERS", "1") + model = Model() + first = torch.tensor([1.0]) + second = torch.tensor([2.0]) + + assert load_weights_maybe_cached(model, [("expert", first)]) == {"expert"} + cache = model._nrl_refit_loader_cache + old_local = model.local + model.local = torch.nn.Parameter(torch.zeros(1), requires_grad=False) + model.local.weight_loader = make_loader("replacement") + + assert load_weights_maybe_cached(model, [("expert", second)]) == {"expert"} + + assert model.load_calls == [["expert"], ["expert"]] + assert events == ["remote", "local", "remote", "replacement"] + torch.testing.assert_close(old_local, first) + torch.testing.assert_close(model.local, second) + assert cache.calls == {} + assert cache.uncached == set() + assert cache.snapshot == {} + + @pytest.mark.vllm def test_refit_load_weights_dispatches_to_sharded_path_when_enabled(): from nemo_rl.models.generation.vllm.vllm_backend import ( diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index 80c7df8364c..17f058c87b8 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -18,6 +18,7 @@ from pathlib import Path from types import SimpleNamespace from typing import Any, Optional +from unittest.mock import MagicMock import numpy as np import pytest @@ -136,6 +137,220 @@ def prepare_refit_info( assert exported[scale_name].shape == (64, 2) +def test_reference_model_pinned_swap_restores_state_and_reuses_buffer(monkeypatch): + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + + class Model(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.tensor([1.0])) + self.register_buffer("extra_state_cache", torch.tensor([2.0])) + + worker = object.__new__(MegatronPolicyWorkerImpl) + worker.model = Model() + worker.reference_state_dict = { + "weight": torch.tensor([11.0]), + "extra_state_cache": torch.tensor([12.0]), + } + worker.cfg = { + "megatron_cfg": { + "pinned_reference_swap": True, + "empty_unused_memory_level": 0, + } + } + worker._pinned_swap_save_buffers = {} + worker.should_disable_forward_pre_hook = False + worker.sampling_params = None + allocations = [] + original_empty = torch.empty + + def empty(*args, **kwargs): + allocations.append(kwargs) + kwargs = {**kwargs, "pin_memory": False} + return original_empty(*args, **kwargs) + + synchronize = MagicMock() + monkeypatch.setattr(torch, "empty", empty) + monkeypatch.setattr(torch.cuda, "synchronize", synchronize) + + cached_buffer = None + for _ in range(2): + with worker.use_reference_model(): + torch.testing.assert_close(worker.model.weight, torch.tensor([11.0])) + torch.testing.assert_close( + worker.model.extra_state_cache, torch.tensor([12.0]) + ) + torch.testing.assert_close(worker.model.weight, torch.tensor([1.0])) + torch.testing.assert_close(worker.model.extra_state_cache, torch.tensor([2.0])) + if cached_buffer is None: + cached_buffer = worker._pinned_swap_save_buffers["weight"] + else: + assert worker._pinned_swap_save_buffers["weight"] is cached_buffer + + assert list(worker._pinned_swap_save_buffers) == ["weight"] + assert allocations == [ + { + "dtype": torch.float32, + "device": "cpu", + "pin_memory": True, + } + ] + assert synchronize.call_count == 6 + + +def test_clear_rope_and_moe_dispatcher_caches_clears_tensor_state(monkeypatch): + from megatron.core.models.common.embeddings import rotary_pos_embedding + + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + + cache_clear = MagicMock() + + def forward(): + return None + + forward.cache_clear = cache_clear + monkeypatch.setattr( + rotary_pos_embedding, + "RotaryEmbedding", + SimpleNamespace(forward=forward), + ) + dispatcher = SimpleNamespace( + probs=torch.ones(1), + routing_map=torch.ones(1), + reversed_local_input_permutation_mapping=torch.ones(1), + local_probs=torch.ones(1), + local_map=torch.ones(1), + non_tensor="keep", + ) + worker = object.__new__(MegatronPolicyWorkerImpl) + worker.model = SimpleNamespace( + modules=lambda: [ + SimpleNamespace(), + SimpleNamespace(token_dispatcher=None), + SimpleNamespace(token_dispatcher=dispatcher), + ] + ) + + worker._clear_rope_and_moe_dispatcher_caches() + + cache_clear.assert_called_once_with() + assert dispatcher.probs is None + assert dispatcher.routing_map is None + assert dispatcher.reversed_local_input_permutation_mapping is None + assert dispatcher.local_probs is None + assert dispatcher.local_map is None + assert dispatcher.non_tensor == "keep" + + +def test_clear_rope_and_moe_dispatcher_caches_is_best_effort(monkeypatch): + from megatron.core.models.common.embeddings import rotary_pos_embedding + + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + + def forward(): + return None + + forward.cache_clear = MagicMock(side_effect=RuntimeError("rotary cache")) + monkeypatch.setattr( + rotary_pos_embedding, + "RotaryEmbedding", + SimpleNamespace(forward=forward), + ) + worker = object.__new__(MegatronPolicyWorkerImpl) + worker.model = SimpleNamespace( + modules=MagicMock(side_effect=RuntimeError("module traversal")) + ) + + worker._clear_rope_and_moe_dispatcher_caches() + + +@pytest.mark.parametrize( + ("selected", "dtype"), + [(False, torch.bfloat16), (True, torch.float8_e4m3fn)], +) +def test_maybe_prequantize_param_passthrough(selected, dtype): + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + + worker = object.__new__(MegatronPolicyWorkerImpl) + name = "model.weight" + worker._refit_prequant_names = {name} if selected else set() + tensor = torch.ones(2, 2, dtype=dtype) + + result = list(worker._maybe_prequantize_param(name, tensor)) + + assert len(result) == 1 + assert result[0][0] == name + assert result[0][1] is tensor + + +@pytest.mark.parametrize("slim", [False, True]) +def test_offload_after_refit_routes_cleanup_by_mode(monkeypatch, slim): + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + + worker = object.__new__(MegatronPolicyWorkerImpl) + model = SimpleNamespace(eval=MagicMock()) + worker.model = model + worker.move_model = MagicMock(return_value=model) + worker.cfg = { + "megatron_cfg": { + "refit_slim_offload_after": slim, + "clear_memory_caches_before_refit": True, + } + } + worker.fp8_cfg = {"force_clear_fp8_caches": True} + worker._clear_fp8_caches = MagicMock() + worker._clear_rope_and_moe_dispatcher_caches = MagicMock() + worker.optimizer = object() + worker.optimizer_cpu_offload = False + worker.move_optimizer = MagicMock() + worker.offload_before_refit = MagicMock() + collect = MagicMock() + empty_cache = MagicMock() + monkeypatch.setattr( + torch, + "randn", + lambda *_args, **_kwargs: SimpleNamespace(cuda=lambda: None), + ) + monkeypatch.setattr(torch.cuda, "memory_allocated", lambda: 0) + monkeypatch.setattr(torch.cuda, "memory_reserved", lambda: 0) + monkeypatch.setattr(torch.cuda, "empty_cache", empty_cache) + monkeypatch.setattr(torch.cuda.nvtx, "range_push", lambda *_args: None) + monkeypatch.setattr(torch.cuda.nvtx, "range_pop", lambda: None) + monkeypatch.setattr( + "nemo_rl.models.policy.workers.megatron_policy_worker.gc.collect", + collect, + ) + + worker.offload_after_refit() + + worker.move_model.assert_called_once_with(model, "cpu") + model.eval.assert_called_once_with() + if slim: + worker._clear_fp8_caches.assert_called_once_with() + worker._clear_rope_and_moe_dispatcher_caches.assert_called_once_with() + worker.move_optimizer.assert_called_once_with("cpu") + collect.assert_called_once_with() + empty_cache.assert_called_once_with() + worker.offload_before_refit.assert_not_called() + else: + worker.offload_before_refit.assert_called_once_with() + worker._clear_fp8_caches.assert_not_called() + worker._clear_rope_and_moe_dispatcher_caches.assert_not_called() + worker.move_optimizer.assert_not_called() + collect.assert_not_called() + empty_cache.assert_not_called() + + def test_megatron_prepare_for_training_restores_optimizer(): from nemo_rl.models.policy.workers.megatron_policy_worker import ( MegatronPolicyWorkerImpl, diff --git a/tests/unit/models/policy/test_policy_validation.py b/tests/unit/models/policy/test_policy_validation.py index 6a4c6a8d3ab..67e4c858086 100644 --- a/tests/unit/models/policy/test_policy_validation.py +++ b/tests/unit/models/policy/test_policy_validation.py @@ -35,6 +35,30 @@ def test_shutdown_succeeds_before_worker_group_is_initialized(capsys) -> None: assert capsys.readouterr().out == "" +def test_enable_refit_prequantize_forwards_names_and_returns_metadata( + monkeypatch, +) -> None: + policy = Policy.__new__(Policy) + futures = [object()] + updated_info = { + "model.weight": ((2, 2), "float8_e4m3fn"), + "model.weight_scale_from_checkpoint": ((2, 1), "uint8"), + } + policy.worker_group = MagicMock() + policy.worker_group.run_all_workers_single_data.return_value = futures + ray_get = MagicMock(return_value=[updated_info]) + monkeypatch.setattr("nemo_rl.models.policy.lm_policy.ray.get", ray_get) + + result = policy.enable_refit_prequantize(["model.weight"]) + + assert result is updated_info + policy.worker_group.run_all_workers_single_data.assert_called_once_with( + "enable_refit_prequantize", + param_names=["model.weight"], + ) + ray_get.assert_called_once_with(futures) + + def create_mock_cluster(world_size: int): """Create a mock cluster with the specified world size.""" cluster = MagicMock() diff --git a/tests/unit/weight_sync/test_checkpoint_engine_weight_synchronizer.py b/tests/unit/weight_sync/test_checkpoint_engine_weight_synchronizer.py index ecb8585acb9..3d9d16d66c9 100644 --- a/tests/unit/weight_sync/test_checkpoint_engine_weight_synchronizer.py +++ b/tests/unit/weight_sync/test_checkpoint_engine_weight_synchronizer.py @@ -172,6 +172,23 @@ def test_init_communicator_completes_prequant_handshake(self): ] sync._ensure_checkpoint_engine_ready.assert_called_once_with() + def test_init_communicator_rejects_missing_prequant_metadata(self): + sync = _checkpoint_sync(MagicMock()) + sync._ensure_checkpoint_engine_ready = MagicMock() + sync._policy.prepare_refit_info.return_value = { + "layer_0": {"shape": [4096, 4096]} + } + sync._generation.prepare_refit_info.return_value = ["layer_0"] + sync._policy.enable_refit_prequantize.return_value = None + + with pytest.raises( + RuntimeError, + match="did not return updated metadata", + ): + sync.init_communicator() + + sync._ensure_checkpoint_engine_ready.assert_not_called() + @patch("nemo_rl.weight_sync.checkpoint_engine_weight_synchronizer.ray") def test_bucket_uses_minimum_total_memory_and_is_cached(self, mock_ray, capsys): config = _checkpoint_engine_cfg(bucket_memory_ratio=0.125) From b237fd2554a5baf07fdfac0c553aabcb8d603921 Mon Sep 17 00:00:00 2001 From: sna Date: Mon, 27 Jul 2026 15:52:06 -0700 Subject: [PATCH 06/76] test(refit): fix vLLM test import order Signed-off-by: sna --- tests/unit/models/generation/test_vllm_fp8_quantization.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index cdc5589f97c..b349e8f7493 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -175,9 +175,10 @@ def fake_patch(path, _replacement): def test_load_weights_preserves_prequantized_mxfp8_and_clamps_scales( fp8_module, monkeypatch ): - from nemo_rl.models.generation.vllm import vllm_backend from vllm.model_executor.layers.quantization.utils import mxfp8_utils + from nemo_rl.models.generation.vllm import vllm_backend + fp8 = fp8_module fp8.global_fp8_config = types.SimpleNamespace(is_mx=True) native = torch.ones(2, 2, dtype=torch.bfloat16) From ab030c30a840771a4a7b6ac1cc231e65ba9c41e1 Mon Sep 17 00:00:00 2001 From: sna Date: Tue, 28 Jul 2026 12:42:17 -0700 Subject: [PATCH 07/76] test(fp8): cover MXFP8 MoE padding path Signed-off-by: sna --- .../generation/test_vllm_fp8_quantization.py | 142 ++++++++++++++++++ 1 file changed, 142 insertions(+) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index b349e8f7493..07678bf3cae 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -13,6 +13,7 @@ # limitations under the License. import types +from typing import Any import pytest import torch @@ -232,3 +233,144 @@ def test_load_weights_preserves_prequantized_mxfp8_and_clamps_scales( loaded[3][1], torch.tensor([[1, 7], [3, 1]], dtype=torch.uint8), ) + + +def test_mxfp8_padding_helpers_preserve_values_and_fill_padding( + fp8_module: types.ModuleType, +) -> None: + fp8 = fp8_module + tensor = torch.arange(12).reshape(2, 2, 3) + + assert fp8._round_up(1856, 128) == 1920 + assert fp8._pad_tensor_dim(tensor, 1, 2) is tensor + + padded = fp8._pad_tensor_dim(tensor, 1, 4, pad_value=7) + torch.testing.assert_close(padded[:, :2], tensor) + torch.testing.assert_close(padded[:, 2:], torch.full((2, 2, 3), 7)) + + w13 = torch.arange(24).reshape(2, 4, 3) + assert fp8._pad_w13_shards(w13, 2, 2) is w13 + + padded_w13 = fp8._pad_w13_shards(w13, 2, 3, pad_value=9) + expected_w13 = torch.tensor( + [ + [[0, 1, 2], [3, 4, 5], [9, 9, 9], [6, 7, 8], [9, 10, 11], [9, 9, 9]], + [ + [12, 13, 14], + [15, 16, 17], + [9, 9, 9], + [18, 19, 20], + [21, 22, 23], + [9, 9, 9], + ], + ] + ) + torch.testing.assert_close(padded_w13, expected_w13) + torch.testing.assert_close( + fp8._clamp_mxfp8_scale(torch.tensor([0, 2, 0], dtype=torch.uint8)), + torch.tensor([1, 2, 1], dtype=torch.uint8), + ) + + +def test_set_mxfp8_apply_tensor_reuses_matching_storage( + fp8_module: types.ModuleType, +) -> None: + fp8 = fp8_module + layer = torch.nn.Module() + + fp8._set_mxfp8_apply_tensor(layer, "weight_for_apply", torch.ones(2, 3)) + first = layer.weight_for_apply + first_data_ptr = first.data_ptr() + + fp8._set_mxfp8_apply_tensor(layer, "weight_for_apply", torch.full((2, 3), 4.0)) + + assert layer.weight_for_apply is first + assert layer.weight_for_apply.data_ptr() == first_data_ptr + torch.testing.assert_close(layer.weight_for_apply, torch.full((2, 3), 4.0)) + + +def test_process_mxfp8_moe_pads_kernel_tensors_without_changing_checkpoint_layout( + fp8_module: types.ModuleType, + monkeypatch: pytest.MonkeyPatch, +) -> None: + fp8 = fp8_module + captured: dict[str, Any] = {} + + def fake_batched_shuffle( + layer: torch.nn.Module, + w13_weight: torch.Tensor, + w2_weight: torch.Tensor, + w13_scale: torch.Tensor, + w2_scale: torch.Tensor, + is_gated: bool, + epilogue_tile_m: int, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + captured.update( + { + "layer": layer, + "w13_weight": w13_weight, + "w2_weight": w2_weight, + "w13_scale": w13_scale, + "w2_scale": w2_scale, + "is_gated": is_gated, + "epilogue_tile_m": epilogue_tile_m, + } + ) + return tuple( + tensor.clone() for tensor in (w13_weight, w2_weight, w13_scale, w2_scale) + ) + + monkeypatch.setattr(fp8, "_shuffle_mxfp8_moe_batched", fake_batched_shuffle) + monkeypatch.delenv("NRL_MXFP8_BATCHED_SHUFFLE", raising=False) + + layer = torch.nn.Module() + layer.w13_weight = torch.nn.Parameter( + torch.arange(30, dtype=torch.float32).reshape(2, 3, 5), + requires_grad=False, + ) + layer.w2_weight = torch.nn.Parameter( + torch.arange(30, dtype=torch.float32).reshape(2, 5, 3), + requires_grad=False, + ) + layer.w13_weight_scale_from_checkpoint = torch.nn.Parameter( + torch.zeros(2, 3, 1, dtype=torch.uint8), + requires_grad=False, + ) + layer.w2_weight_scale_from_checkpoint = torch.nn.Parameter( + torch.zeros(2, 5, 1, dtype=torch.uint8), + requires_grad=False, + ) + layer.moe_config = types.SimpleNamespace(intermediate_size_per_partition=3) + quant_method = types.SimpleNamespace( + moe=types.SimpleNamespace( + is_act_and_mul=False, + intermediate_size_per_partition=3, + ) + ) + original_w13 = layer.w13_weight.detach().clone() + original_w2 = layer.w2_weight.detach().clone() + + fp8.process_weights_after_loading_mxfp8_moe(quant_method, layer) + + assert captured["layer"] is layer + assert captured["is_gated"] is False + assert captured["epilogue_tile_m"] == 128 + assert captured["w13_weight"].shape == (2, 128, 512) + assert captured["w2_weight"].shape == (2, 512, 128) + assert captured["w13_scale"].shape == (2, 128, 16) + assert captured["w2_scale"].shape == (2, 512, 4) + assert torch.count_nonzero(captured["w13_scale"] == 0) == 0 + assert torch.count_nonzero(captured["w2_scale"] == 0) == 0 + + torch.testing.assert_close(layer.w13_weight, original_w13) + torch.testing.assert_close(layer.w2_weight, original_w2) + assert layer.mxfp8_unpadded_hidden_size == 5 + assert layer.mxfp8_padded_hidden_size == 512 + assert layer.mxfp8_unpadded_intermediate_size_per_partition == 3 + assert layer.mxfp8_padded_intermediate_size_per_partition == 128 + assert layer.moe_config.intermediate_size_per_partition == 128 + assert quant_method.moe.intermediate_size_per_partition == 128 + assert layer.w13_weight_for_apply.shape == (2, 128, 512) + assert layer.w2_weight_for_apply.shape == (2, 512, 128) + assert layer.w13_scale_for_apply.shape == (2, 128, 16) + assert layer.w2_scale_for_apply.shape == (2, 512, 4) From f391a4a7c1aa454ad6b45d2fc9dbd2fc021be28e Mon Sep 17 00:00:00 2001 From: seonjinn Date: Wed, 29 Jul 2026 00:22:05 -0700 Subject: [PATCH 08/76] fix(refit): reconcile MXFP8 optimizations with NCCL reshard Signed-off-by: seonjinn --- .../models/generation/vllm/vllm_backend.py | 3 -- nemo_rl/weight_sync/nccl_reshard_utils.py | 7 ++++ .../models/generation/test_vllm_backend.py | 35 +++++++++++++++++++ 3 files changed, 42 insertions(+), 3 deletions(-) diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index 31a9f7f2d0c..2135c6cb504 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -1199,9 +1199,6 @@ def _recv_one_param(param_info, group, stream): ) torch.cuda.empty_cache() - - # Finalize FP8 KV-cache per-layer k/v scales after the misc broadcast. - self._maybe_process_fp8_kv_cache() return True def _receive_and_load_misc_params(self) -> None: diff --git a/nemo_rl/weight_sync/nccl_reshard_utils.py b/nemo_rl/weight_sync/nccl_reshard_utils.py index 8dd978e6772..af1df04b767 100644 --- a/nemo_rl/weight_sync/nccl_reshard_utils.py +++ b/nemo_rl/weight_sync/nccl_reshard_utils.py @@ -561,6 +561,13 @@ def check_nccl_reshard_refit_support(master_config: dict) -> None: "dynamic expert load balancing can change ownership afterwards)." ) + if vllm_cfg.get("refit_prequantize", False): + violations.append( + "policy.generation.vllm_cfg.refit_prequantize must be False " + "(nccl_reshard_refit requires matching training and generation " + "storage and does not use BF16-to-MXFP8 prequantization)." + ) + # This initial version supports only the Megatron train + vLLM gen # combination; the DTensor train backend refit path is intentionally # dropped (Megatron is the path validated end-to-end). diff --git a/tests/unit/models/generation/test_vllm_backend.py b/tests/unit/models/generation/test_vllm_backend.py index 4b892ef178d..cb8b5a84f14 100644 --- a/tests/unit/models/generation/test_vllm_backend.py +++ b/tests/unit/models/generation/test_vllm_backend.py @@ -299,6 +299,41 @@ def packed_broadcast_consumer(iterator, group, src, post_unpack_func): assert call_order == expected_call_order +@pytest.mark.vllm +def test_nccl_reshard_refit_processes_weights_once(monkeypatch): + from nemo_rl.models.generation.vllm import vllm_backend + + ext = vllm_backend.VllmInternalWorkerExtension.__new__( + vllm_backend.VllmInternalWorkerExtension + ) + ext.nccl_reshard_refit_info = { + "layer_names": [], + "per_layer_params": {}, + } + ext.model_runner = SimpleNamespace(model=object(), vllm_config=object()) + ext.model_config = object() + ext.device = object() + ext._receive_and_load_misc_params = MagicMock() + + monkeypatch.setattr(vllm_backend.torch.cuda, "Stream", MagicMock) + monkeypatch.setattr(vllm_backend.torch.cuda, "synchronize", MagicMock()) + monkeypatch.setattr(vllm_backend.torch.cuda, "empty_cache", MagicMock()) + monkeypatch.setattr(vllm_backend.torch.distributed, "get_rank", lambda: 1) + monkeypatch.setattr( + "vllm.config.set_current_vllm_config", lambda _: contextlib.nullcontext() + ) + process_weights = MagicMock() + monkeypatch.setattr( + "vllm.model_executor.model_loader.utils.process_weights_after_loading", + process_weights, + ) + + assert ext.nccl_reshard_refit() is True + process_weights.assert_called_once_with( + ext.model_runner.model, ext.model_config, ext.device + ) + + @pytest.mark.vllm @pytest.mark.parametrize( "method_name", From 85a6feaf581cd8875180b40a6bd52b53938bdf6f Mon Sep 17 00:00:00 2001 From: seonjinn Date: Wed, 29 Jul 2026 10:22:23 -0700 Subject: [PATCH 09/76] fix(refit): validate checkpoint prequant backend Signed-off-by: seonjinn --- .../checkpoint_engine_weight_synchronizer.py | 7 +++++++ tests/unit/models/policy/test_megatron_worker.py | 2 +- .../test_checkpoint_engine_weight_synchronizer.py | 13 +++++++++++++ 3 files changed, 21 insertions(+), 1 deletion(-) diff --git a/nemo_rl/weight_sync/checkpoint_engine_weight_synchronizer.py b/nemo_rl/weight_sync/checkpoint_engine_weight_synchronizer.py index 76af722cb62..edf12b0eca6 100644 --- a/nemo_rl/weight_sync/checkpoint_engine_weight_synchronizer.py +++ b/nemo_rl/weight_sync/checkpoint_engine_weight_synchronizer.py @@ -70,6 +70,13 @@ def init_communicator(self) -> None: state_dict_info = self._policy.prepare_refit_info() prequant_names = self._generation.prepare_refit_info(state_dict_info) if prequant_names: + megatron_cfg = self._policy.cfg.get("megatron_cfg", {}) + if not megatron_cfg.get("enabled", False): + raise ValueError( + "vllm_cfg.refit_prequantize requires the Megatron policy backend " + "(policy.megatron_cfg.enabled=true); the DTensor workers do not " + "implement trainer-side pre-quantized refit." + ) updated_info = self._policy.enable_refit_prequantize(prequant_names) if updated_info is None: raise RuntimeError( diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index 17f058c87b8..d471a2f40a8 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -104,7 +104,7 @@ class _PrequantCheckpointWorker(MegatronCheckpointEngineSendMixin): worker.model = object() worker.draft_model = None worker.refit_conversion_tasks = [] - worker.cfg = {} + worker.cfg = {"megatron_cfg": {"enabled": True}} worker.megatron_bridge = SimpleNamespace( export_hf_weights=lambda *_args, **_kwargs: iter([(name, weight)]) ) diff --git a/tests/unit/weight_sync/test_checkpoint_engine_weight_synchronizer.py b/tests/unit/weight_sync/test_checkpoint_engine_weight_synchronizer.py index 3d9d16d66c9..8a547175ec5 100644 --- a/tests/unit/weight_sync/test_checkpoint_engine_weight_synchronizer.py +++ b/tests/unit/weight_sync/test_checkpoint_engine_weight_synchronizer.py @@ -33,6 +33,7 @@ def _mock_policy(**overrides): policy = MagicMock() + policy.cfg = {"megatron_cfg": {"enabled": True}} policy.offload_before_refit.return_value = None policy.offload_after_refit.return_value = None policy.prepare_refit_info.return_value = {"layer_0": {"shape": [4096, 4096]}} @@ -189,6 +190,18 @@ def test_init_communicator_rejects_missing_prequant_metadata(self): sync._ensure_checkpoint_engine_ready.assert_not_called() + def test_init_communicator_rejects_prequantization_without_megatron(self): + sync = _checkpoint_sync(MagicMock()) + sync._ensure_checkpoint_engine_ready = MagicMock() + sync._policy.cfg = {"megatron_cfg": {"enabled": False}} + sync._generation.prepare_refit_info.return_value = ["layer_0"] + + with pytest.raises(ValueError, match="requires the Megatron policy backend"): + sync.init_communicator() + + sync._policy.enable_refit_prequantize.assert_not_called() + sync._ensure_checkpoint_engine_ready.assert_not_called() + @patch("nemo_rl.weight_sync.checkpoint_engine_weight_synchronizer.ray") def test_bucket_uses_minimum_total_memory_and_is_cached(self, mock_ray, capsys): config = _checkpoint_engine_cfg(bucket_memory_ratio=0.125) From 716bb7480802e5f697ad4769f1fab797021f1f03 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Wed, 29 Jul 2026 10:32:26 -0700 Subject: [PATCH 10/76] fix(refit): reject MXFP8 with NCCL reshard Signed-off-by: seonjinn --- nemo_rl/weight_sync/nccl_reshard_utils.py | 6 ++++ .../weight_sync/test_nccl_reshard_utils.py | 36 +++++++++++++++++++ 2 files changed, 42 insertions(+) diff --git a/nemo_rl/weight_sync/nccl_reshard_utils.py b/nemo_rl/weight_sync/nccl_reshard_utils.py index af1df04b767..e11e3a5cdb4 100644 --- a/nemo_rl/weight_sync/nccl_reshard_utils.py +++ b/nemo_rl/weight_sync/nccl_reshard_utils.py @@ -568,6 +568,12 @@ def check_nccl_reshard_refit_support(master_config: dict) -> None: "storage and does not use BF16-to-MXFP8 prequantization)." ) + if vllm_cfg.get("is_mx", False): + violations.append( + "policy.generation.vllm_cfg.is_mx must be False " + "(nccl_reshard_refit does not support MXFP8 storage or scale mapping)." + ) + # This initial version supports only the Megatron train + vLLM gen # combination; the DTensor train backend refit path is intentionally # dropped (Megatron is the path validated end-to-end). diff --git a/tests/unit/weight_sync/test_nccl_reshard_utils.py b/tests/unit/weight_sync/test_nccl_reshard_utils.py index e8d2fff269f..18ee1edc8a6 100644 --- a/tests/unit/weight_sync/test_nccl_reshard_utils.py +++ b/tests/unit/weight_sync/test_nccl_reshard_utils.py @@ -87,6 +87,42 @@ def test_check_nccl_reshard_refit_support_rejects_invalid_config( assert expected_violation in str(exc_info.value) +@pytest.mark.parametrize( + ("vllm_cfg_update", "expected_violation"), + [ + ( + {"refit_prequantize": True}, + "policy.generation.vllm_cfg.refit_prequantize must be False", + ), + ( + {"is_mx": True}, + "policy.generation.vllm_cfg.is_mx must be False", + ), + ], +) +def test_check_nccl_reshard_refit_support_rejects_unsupported_refit_modes( + vllm_cfg_update: dict[str, object], expected_violation: str +) -> None: + config = _valid_nccl_reshard_config() + config.policy["generation"]["vllm_cfg"].update(vllm_cfg_update) + + with pytest.raises(ValueError) as exc_info: + check_nccl_reshard_refit_support(config) + + assert expected_violation in str(exc_info.value) + + +def test_check_nccl_reshard_refit_support_accepts_matching_blockwise_fp8() -> None: + config = _valid_nccl_reshard_config() + config.policy["generation"]["vllm_cfg"]["precision"] = "fp8" + config.policy["megatron_cfg"]["fp8_cfg"] = { + "fp8_param": True, + "fp8_recipe": "blockwise", + } + + check_nccl_reshard_refit_support(config) + + # -------------------------------------------------------------------------- # MeshInfo # -------------------------------------------------------------------------- From f638fa6a313e91b3d960b42aafceeb9eea8dd43a Mon Sep 17 00:00:00 2001 From: seonjinn Date: Wed, 29 Jul 2026 10:55:56 -0700 Subject: [PATCH 11/76] fix(refit): validate NCCL reshard storage precision Signed-off-by: seonjinn --- nemo_rl/weight_sync/nccl_reshard_utils.py | 21 +++++++++++++-- .../weight_sync/test_nccl_reshard_utils.py | 26 +++++++++++++++++++ 2 files changed, 45 insertions(+), 2 deletions(-) diff --git a/nemo_rl/weight_sync/nccl_reshard_utils.py b/nemo_rl/weight_sync/nccl_reshard_utils.py index e11e3a5cdb4..9bfb4c9dfb6 100644 --- a/nemo_rl/weight_sync/nccl_reshard_utils.py +++ b/nemo_rl/weight_sync/nccl_reshard_utils.py @@ -535,6 +535,7 @@ def check_nccl_reshard_refit_support(master_config: dict) -> None: generation = policy.get("generation", {}) or {} megatron_cfg = policy.get("megatron_cfg", {}) or {} dtensor_cfg = policy.get("dtensor_cfg", {}) or {} + policy_precision = policy.get("precision") vllm_cfg = generation.get("vllm_cfg", {}) or {} vllm_kwargs = generation.get("vllm_kwargs", {}) or {} @@ -554,6 +555,13 @@ def check_nccl_reshard_refit_support(master_config: dict) -> None: f"policy.generation.backend must be 'vllm' (got {backend!r})." ) + if policy_precision != "bfloat16": + violations.append( + "policy.precision must be 'bfloat16' for nccl_reshard_refit " + f"(got {policy_precision!r}); the refit byte-copies training storage " + "into the generation model." + ) + if vllm_kwargs.get("enable_eplb"): violations.append( "policy.generation.vllm_kwargs.enable_eplb must be False " @@ -626,10 +634,19 @@ def check_nccl_reshard_refit_support(master_config: dict) -> None: # BF16→FP8 (train-side quant on the fly) is not implemented; FP8→BF16 # has no consumer (vLLM doesn't accept FP8 bytes into a BF16 param). fp8_cfg = megatron_cfg.get("fp8_cfg", {}) or {} - fp8_param = fp8_cfg.get("fp8_param", False) + fp8_enabled = bool(fp8_cfg.get("enabled", False)) + fp8_param_requested = bool(fp8_cfg.get("fp8_param", False)) + fp8_param = fp8_enabled and fp8_param_requested fp8_recipe = fp8_cfg.get("fp8_recipe", None) gen_precision = vllm_cfg.get("precision", None) + if fp8_param_requested and not fp8_enabled: + violations.append( + "policy.megatron_cfg.fp8_cfg.fp8_param=True requires " + "policy.megatron_cfg.fp8_cfg.enabled=True; disabled FP8 does not " + "produce FP8 parameter storage for nccl_reshard_refit." + ) + # The refit byte-copies weights train -> gen, so gen dtype must match # train: BF16 (unset / "auto" / "bf16" / "bfloat16") or FP8 ("fp8"). A # value like "float16"/"float32" would silently mismatch the bf16 train @@ -648,7 +665,7 @@ def check_nccl_reshard_refit_support(master_config: dict) -> None: if not fp8_param: violations.append( "policy.generation.vllm_cfg.precision='fp8' requires " - "policy.megatron_cfg.fp8_cfg.fp8_param=True " + "policy.megatron_cfg.fp8_cfg.enabled=True and fp8_param=True " "(BF16→FP8 train-side quantization is not implemented yet)." ) elif fp8_recipe != "blockwise": diff --git a/tests/unit/weight_sync/test_nccl_reshard_utils.py b/tests/unit/weight_sync/test_nccl_reshard_utils.py index 18ee1edc8a6..0f96668211f 100644 --- a/tests/unit/weight_sync/test_nccl_reshard_utils.py +++ b/tests/unit/weight_sync/test_nccl_reshard_utils.py @@ -47,6 +47,7 @@ def _valid_nccl_reshard_config() -> SimpleNamespace: return SimpleNamespace( policy={ + "precision": "bfloat16", "generation": { "backend": "vllm", "colocated": {"enabled": False}, @@ -116,6 +117,7 @@ def test_check_nccl_reshard_refit_support_accepts_matching_blockwise_fp8() -> No config = _valid_nccl_reshard_config() config.policy["generation"]["vllm_cfg"]["precision"] = "fp8" config.policy["megatron_cfg"]["fp8_cfg"] = { + "enabled": True, "fp8_param": True, "fp8_recipe": "blockwise", } @@ -123,6 +125,30 @@ def test_check_nccl_reshard_refit_support_accepts_matching_blockwise_fp8() -> No check_nccl_reshard_refit_support(config) +def test_check_nccl_reshard_refit_support_rejects_disabled_fp8_param_storage() -> None: + config = _valid_nccl_reshard_config() + config.policy["generation"]["vllm_cfg"]["precision"] = "fp8" + config.policy["megatron_cfg"]["fp8_cfg"] = { + "enabled": False, + "fp8_param": True, + "fp8_recipe": "blockwise", + } + + with pytest.raises(ValueError, match="fp8_cfg.enabled=True"): + check_nccl_reshard_refit_support(config) + + +@pytest.mark.parametrize("precision", ["float16", "float32"]) +def test_check_nccl_reshard_refit_support_rejects_non_bfloat16_policy_precision( + precision: str, +) -> None: + config = _valid_nccl_reshard_config() + config.policy["precision"] = precision + + with pytest.raises(ValueError, match="policy.precision must be 'bfloat16'"): + check_nccl_reshard_refit_support(config) + + # -------------------------------------------------------------------------- # MeshInfo # -------------------------------------------------------------------------- From 2fde3b08c60a1c31145d354b654227fb22e3cd06 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Wed, 29 Jul 2026 22:56:23 -0700 Subject: [PATCH 12/76] ci: associate Codecov uploads with pull requests Signed-off-by: seonjinn --- .github/workflows/cicd-main.yml | 1 + 1 file changed, 1 insertion(+) diff --git a/.github/workflows/cicd-main.yml b/.github/workflows/cicd-main.yml index 93dec1b806f..5f0f5030041 100644 --- a/.github/workflows/cicd-main.yml +++ b/.github/workflows/cicd-main.yml @@ -1102,6 +1102,7 @@ jobs: verbose: true flags: ${{ matrix.flag }} base_sha: ${{ fromJSON(steps.get-pr-info.outputs.pr-info || '{}').base.sha }} + override_pr: ${{ fromJSON(steps.get-pr-info.outputs.pr-info || '{}').number }} - name: Upload artifacts if: ${{ steps.check-artifacts.outputs.artifacts-found == 'true' }} From 00fa128513dc55e13201ad0f90f17876fe5f97f4 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Thu, 30 Jul 2026 10:20:36 -0700 Subject: [PATCH 13/76] fix: address MXFP8 refit review feedback Signed-off-by: seonjinn --- docs/fp8.md | 19 +++ docs/guides/refit.md | 9 ++ examples/configs/grpo_math_1B.yaml | 2 + .../grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml | 1 + ...3-235b-32n4g-async-1off-mxfp8-rollout.yaml | 1 + ...-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml | 1 + .../grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml | 1 + .../grpo-qwen3-32b-4n4g-mxfp8-rollout.yaml | 1 + ...en3-32b-8n4g-async-1off-mxfp8-rollout.yaml | 1 + nemo_rl/algorithms/distillation.py | 8 +- nemo_rl/algorithms/grpo.py | 9 +- nemo_rl/algorithms/ppo.py | 9 +- nemo_rl/algorithms/utils.py | 37 +---- nemo_rl/models/generation/__init__.py | 6 +- nemo_rl/models/generation/vllm/config.py | 18 +++ .../generation/vllm/quantization/fp8.py | 79 +++++++---- .../vllm/quantization/fp8_train_utils.py | 11 +- .../models/generation/vllm/vllm_generation.py | 7 +- nemo_rl/models/policy/interfaces.py | 7 +- .../checkpoint_engine_weight_synchronizer.py | 22 +-- .../collective_weight_synchronizer.py | 8 +- .../weight_sync/http_weight_synchronizer.py | 8 +- nemo_rl/weight_sync/interfaces.py | 31 ++++- .../weight_sync/ipc_weight_synchronizer.py | 8 +- .../L1_Functional_Tests_GB200_MXFP8.sh | 2 + .../grpo_vllm_mxfp8_rollout_gb200.sh | 4 + tests/test_mxfp8_rollout_recipes.py | 2 + tests/unit/algorithms/test_utils.py | 63 --------- .../models/generation/test_mxfp8_prequant.py | 64 ++++++++- .../models/generation/test_vllm_config.py | 109 +++++++++++++++ .../unit/reference_configs/grpo_math_1B.yaml | 2 + ..._vllm_remote_sparse_weight_synchronizer.py | 10 ++ .../weight_sync/test_weight_synchronizer.py | 128 +++++++++++++++++- 33 files changed, 508 insertions(+), 180 deletions(-) create mode 100644 tests/unit/models/generation/test_vllm_config.py diff --git a/docs/fp8.md b/docs/fp8.md index da698c950f6..b702f2a2d01 100644 --- a/docs/fp8.md +++ b/docs/fp8.md @@ -57,6 +57,25 @@ FP8 generations are recommended to be configured with the following settings: pow2_activation_scaling_factors: False ``` +For MXFP8 rollout with Megatron training, trainer-side prequantization can +reduce the refit payload: + +```yaml +policy: + generation: + vllm_cfg: + precision: fp8 + is_mx: true + refit_prequantize: true +``` + +`refit_prequantize` is an MXFP8 refit optimization. It requires +`precision: fp8`, `is_mx: true`, and the Megatron policy backend. It moves +eligible weight quantization to the trainer and transfers E4M3 values plus E8M0 +scales instead of BF16 weights. It is rejected for blockwise FP8, BF16, NVFP4, +sparse-delta refit, and NCCL-Reshard refit. NVFP4 real-quant rollout uses its own +packed-weight refit protocol. + To train with FP8, you need to set the Megatron path and configure it using the following settings: ``` diff --git a/docs/guides/refit.md b/docs/guides/refit.md index 86f3d66567d..3abcf974100 100644 --- a/docs/guides/refit.md +++ b/docs/guides/refit.md @@ -35,6 +35,15 @@ the generation backend; both Megatron and DTensor policy workers can send weights. Sparse delta is currently limited to GRPO. NIXL is initialized by the GRPO and distillation setup paths; PPO currently requires colocated generation. +## Performance Options + +| Option | Scope | Effect | +|---|---|---| +| `policy.generation.vllm_cfg.refit_prequantize` | Megatron training with MXFP8 vLLM rollout | Quantizes eligible weights on the trainer and transfers E4M3 values plus E8M0 scales. Requires `precision: fp8` and `is_mx: true`; sparse delta and NCCL Reshard do not support it. | +| `policy.refit_persistent_ipc_buffers` | Colocated CUDA-IPC refit | Reuses the two trainer staging buffers across refits. A fixed `refit_buffer_size_gb` gives stable memory use. | +| `policy.megatron_cfg.refit_slim_offload_after` | Colocated Megatron refit | Avoids repeating grad-buffer offload and a second allocator cleanup after weights are transferred. | +| `policy.megatron_cfg.pinned_reference_swap` | Megatron reference-policy logprobs | Keeps the CPU reference copy in pinned memory for faster host-to-device swaps, at the cost of additional pinned host memory. | + ## Minimal Configuration Colocated refit needs no transport configuration: diff --git a/examples/configs/grpo_math_1B.yaml b/examples/configs/grpo_math_1B.yaml index 6fde615b98e..38cb695e68a 100644 --- a/examples/configs/grpo_math_1B.yaml +++ b/examples/configs/grpo_math_1B.yaml @@ -385,6 +385,8 @@ policy: vllm_cfg: async_engine: false precision: ${policy.precision} + # MXFP8 + Megatron only: quantize on trainer and stream E4M3 values plus scales. + refit_prequantize: false kv_cache_dtype: "auto" tensor_parallel_size: 1 pipeline_parallel_size: 1 diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml index a51d8ddc9f8..53c313db9a6 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml @@ -7,6 +7,7 @@ policy: tensor_parallel_size: 4 precision: "fp8" is_mx: true + refit_prequantize: true quantization_ignored_layer_kws: - q_proj - k_proj diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml index 5f6955bfb17..303e543f1e8 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml @@ -7,6 +7,7 @@ policy: tensor_parallel_size: 4 precision: "fp8" is_mx: true + refit_prequantize: true quantization_ignored_layer_kws: - q_proj - k_proj diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml index 2b160a13829..8d1713d7f15 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml @@ -21,6 +21,7 @@ policy: gpu_memory_utilization: 0.8 precision: "fp8" is_mx: true + refit_prequantize: true quantization_ignored_layer_kws: - q_proj - k_proj diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml index b9fe559bca1..dc4eafefe9b 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml @@ -40,6 +40,7 @@ policy: tensor_parallel_size: 1 precision: "fp8" is_mx: true + refit_prequantize: true quantization_ignored_layer_kws: - q_proj - k_proj diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-mxfp8-rollout.yaml index de338eafe8f..3107e8ff66f 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-mxfp8-rollout.yaml @@ -8,6 +8,7 @@ policy: vllm_cfg: precision: "fp8" is_mx: true + refit_prequantize: true quantization_ignored_layer_kws: - q_proj - k_proj diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-8n4g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-8n4g-async-1off-mxfp8-rollout.yaml index 89555c92fe3..204c96a5dc3 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-8n4g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-8n4g-async-1off-mxfp8-rollout.yaml @@ -9,6 +9,7 @@ policy: gpu_memory_utilization: 0.8 precision: "fp8" is_mx: true + refit_prequantize: true quantization_ignored_layer_kws: - q_proj - k_proj diff --git a/nemo_rl/algorithms/distillation.py b/nemo_rl/algorithms/distillation.py index b33ba118844..faebc9bb58d 100644 --- a/nemo_rl/algorithms/distillation.py +++ b/nemo_rl/algorithms/distillation.py @@ -37,7 +37,7 @@ DistillationLossDataDict, DistillationLossFn, ) -from nemo_rl.algorithms.utils import maybe_enable_refit_prequantize, set_seed +from nemo_rl.algorithms.utils import set_seed from nemo_rl.data import DataConfig from nemo_rl.data.collate_fn import rl_collate_fn from nemo_rl.data.datasets import AllTaskProcessedDataset @@ -90,6 +90,7 @@ checkpoint_engine_refit_config, ) from nemo_rl.weight_sync.factory import create_weight_synchronizer +from nemo_rl.weight_sync.interfaces import initialize_refit_metadata # =============================================================================== # Configuration @@ -630,10 +631,7 @@ def init_nemo_gym(): ) student_generation.weight_synchronizer.init_communicator() elif student_generation is not None: - state_dict_info = student_policy.prepare_refit_info() - maybe_enable_refit_prequantize( - student_policy, student_generation, state_dict_info, master_config.policy - ) + initialize_refit_metadata(student_policy, student_generation) # if it is not colocated inference, initialize collective communication for update weights if not colocated_inference and checkpoint_engine_config is None: diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index bc935e775d9..77d51bd6d9b 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -54,7 +54,6 @@ calculate_baseline_and_std_per_prompt, get_gdpo_reward_component_keys, log_generation_metrics_to_wandb, - maybe_enable_refit_prequantize, print_efficiency_summary, print_performance_metrics, set_seed, @@ -128,6 +127,7 @@ checkpoint_engine_refit_config, ) from nemo_rl.weight_sync.factory import create_weight_synchronizer +from nemo_rl.weight_sync.interfaces import initialize_refit_metadata # =============================================================================== # Configuration @@ -1432,11 +1432,10 @@ def init_trtllm(): ) else: if not (nccl_reshard_refit_enabled and not colocated_inference): - state_dict_info = policy.prepare_refit_info() if policy_generation is not None: - maybe_enable_refit_prequantize( - policy, policy_generation, state_dict_info, master_config.policy - ) + initialize_refit_metadata(policy, policy_generation) + else: + policy.prepare_refit_info() # Spin up non-colocated OPD teacher worker groups AFTER policy / vLLM are # ready. Parallelizing with policy init races on Megatron-Bridge's HF->mcore diff --git a/nemo_rl/algorithms/ppo.py b/nemo_rl/algorithms/ppo.py index 0b772dbc597..a1adc4335f7 100644 --- a/nemo_rl/algorithms/ppo.py +++ b/nemo_rl/algorithms/ppo.py @@ -48,7 +48,6 @@ apply_reward_shaping, ) from nemo_rl.algorithms.utils import ( - maybe_enable_refit_prequantize, print_performance_metrics, set_seed, ) @@ -92,6 +91,7 @@ from nemo_rl.utils.memory_tracker import MemoryTracker from nemo_rl.utils.nsys import maybe_gpu_profile_step from nemo_rl.utils.timer import TimeoutChecker, Timer +from nemo_rl.weight_sync.interfaces import initialize_refit_metadata # =============================================================================== # Configuration @@ -660,11 +660,10 @@ def initialize_generation_with_policy( policy.prepare_for_training() # prepare refit info - state_dict_info = policy.prepare_refit_info() if policy_generation is not None: - maybe_enable_refit_prequantize( - policy, policy_generation, state_dict_info, master_config.policy - ) + initialize_refit_metadata(policy, policy_generation) + else: + policy.prepare_refit_info() # Calculate total setup time total_setup_time = time.perf_counter() - setup_start_time diff --git a/nemo_rl/algorithms/utils.py b/nemo_rl/algorithms/utils.py index 8af9997e3b1..f02485c88dd 100644 --- a/nemo_rl/algorithms/utils.py +++ b/nemo_rl/algorithms/utils.py @@ -16,7 +16,7 @@ import random import warnings from functools import partial, wraps -from typing import TYPE_CHECKING, Any, Optional +from typing import Any, Optional import numpy as np import torch @@ -27,16 +27,10 @@ ) from nemo_rl.data.chat_templates import COMMON_CHAT_TEMPLATES -from nemo_rl.models.policy import PolicyConfig, TokenizerConfig +from nemo_rl.models.policy import TokenizerConfig from nemo_rl.utils.fastokens import maybe_patch_fastokens from nemo_rl.utils.logger import Logger -if TYPE_CHECKING: - # Runtime import would cycle: policy.interfaces pulls in algorithms.loss, - # which imports back into algorithms.utils. - from nemo_rl.models.generation.interfaces import GenerationInterface - from nemo_rl.models.policy.interfaces import ColocatablePolicyInterface - def get_gdpo_reward_component_keys(batch) -> list[str]: """Return batch keys that are named reward components (e.g. reward/correctness) in sorted order.""" @@ -1038,30 +1032,3 @@ def print_efficiency_summary( loggable["efficiency/total_wall_time_s"] = total_wall_time_s return loggable - - -def maybe_enable_refit_prequantize( - policy: "ColocatablePolicyInterface", - policy_generation: "GenerationInterface", - state_dict_info: Optional[dict[str, Any]], - policy_config: PolicyConfig, -) -> None: - """Complete the trainer-side pre-quantized refit handshake if requested. - - The generation backend's prepare_refit_info returns the fp8-eligible - parameter names when vllm_cfg.refit_prequantize is enabled; the trainer - then quantizes exactly those during refit and the receiver's unpack - metadata is refreshed to the quantized dtypes plus scale entries. - """ - prequant_names = policy_generation.prepare_refit_info(state_dict_info) - if not prequant_names: - return - megatron_cfg = policy_config.get("megatron_cfg") - if not (megatron_cfg and megatron_cfg["enabled"]): - raise ValueError( - "vllm_cfg.refit_prequantize requires the Megatron policy backend " - "(policy.megatron_cfg.enabled=true); the DTensor workers do not " - "implement trainer-side pre-quantized refit." - ) - updated_info = policy.enable_refit_prequantize(prequant_names) - policy_generation.prepare_refit_info(updated_info) diff --git a/nemo_rl/models/generation/__init__.py b/nemo_rl/models/generation/__init__.py index 91139aacd40..7c6026b2f47 100644 --- a/nemo_rl/models/generation/__init__.py +++ b/nemo_rl/models/generation/__init__.py @@ -19,7 +19,10 @@ from nemo_rl.models.generation.interfaces import GenerationConfig from nemo_rl.models.generation.trtllm import TrtllmConfig from nemo_rl.models.generation.vllm import VllmConfig -from nemo_rl.models.generation.vllm.config import VLLM_SPARSE_REFIT_TRANSPORTS +from nemo_rl.models.generation.vllm.config import ( + VLLM_SPARSE_REFIT_TRANSPORTS, + validate_vllm_quantization_config, +) TokenizerType = PreTrainedTokenizerBase @@ -46,6 +49,7 @@ def configure_generation_config( # vllm setting if config["backend"] == "vllm": config = cast(VllmConfig, config) + validate_vllm_quantization_config(config) if config.get("real_quant"): export_cpu_offload = config.get("real_quant_export_cpu_offload") if not isinstance(export_cpu_offload, bool): diff --git a/nemo_rl/models/generation/vllm/config.py b/nemo_rl/models/generation/vllm/config.py index 047504e2403..150daa6fd85 100644 --- a/nemo_rl/models/generation/vllm/config.py +++ b/nemo_rl/models/generation/vllm/config.py @@ -158,8 +158,26 @@ class VllmConfig(GenerationConfig): real_quant_ignore: NotRequired[list[str]] +def validate_vllm_quantization_config(config: VllmConfig) -> None: + """Reject quantization options that would otherwise be silently ignored.""" + vllm_cfg = config["vllm_cfg"] + refit_prequantize = vllm_cfg.get("refit_prequantize") + if refit_prequantize is not None and not isinstance(refit_prequantize, bool): + raise ValueError( + "policy.generation.vllm_cfg.refit_prequantize must be a boolean." + ) + if refit_prequantize and not ( + vllm_cfg.get("precision") == "fp8" and vllm_cfg.get("is_mx") is True + ): + raise ValueError( + "policy.generation.vllm_cfg.refit_prequantize requires " + "precision='fp8' and is_mx=true." + ) + + def normalize_vllm_refit_config(config: VllmConfig) -> VllmRefitConfig | None: """Validate the selected refit transport and resolve its scoped defaults.""" + validate_vllm_quantization_config(config) if cast(dict[str, Any], config).get("checkpoint_engine") is not None: raise ValueError( "policy.generation.checkpoint_engine was replaced by " diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index 8dd9ede0ccf..9252763a3e7 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -13,6 +13,7 @@ # limitations under the License. import os +import weakref from dataclasses import dataclass, field from unittest.mock import patch @@ -926,9 +927,9 @@ def _set_mxfp8_apply_tensor(layer, name: str, value: torch.Tensor) -> None: tuple[str, tuple[int, ...], torch.device], torch.Tensor ] = {} -# One-shot flag for NRL_MXFP8_SHUFFLE_VERIFY: compare the batched shuffle -# against the per-expert reference on the first processed layer only. -mxfp8_shuffle_verified = False +# Layers already checked by NRL_MXFP8_SHUFFLE_VERIFY. Weak references avoid +# retaining model modules after a worker shuts down. +mxfp8_shuffle_verified_layers: weakref.WeakSet[object] = weakref.WeakSet() def _mxfp8_scratch(tag: str, shape: torch.Size, device: torch.device) -> torch.Tensor: @@ -1109,6 +1110,38 @@ def _shuffle_mxfp8_moe_per_expert( ) +def _verify_mxfp8_moe_shuffle( + layer, + w13_weight: torch.Tensor, + w2_weight: torch.Tensor, + w13_scale: torch.Tensor, + w2_scale: torch.Tensor, + is_gated: bool, + epilogue_tile_m: int, + batched: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor], +) -> None: + if ( + os.getenv("NRL_MXFP8_SHUFFLE_VERIFY") != "1" + or layer in mxfp8_shuffle_verified_layers + ): + return + + reference = _shuffle_mxfp8_moe_per_expert( + w13_weight, w2_weight, w13_scale, w2_scale, is_gated, epilogue_tile_m + ) + for got, want, tensor_name in zip( + batched, reference, ("w13_weight", "w2_weight", "w13_scale", "w2_scale") + ): + assert torch.equal(got.view(torch.uint8), want.view(torch.uint8)), ( + f"Batched MXFP8 shuffle mismatch vs per-expert reference: {tensor_name}" + ) + mxfp8_shuffle_verified_layers.add(layer) + print( + "[NRL_MXFP8_SHUFFLE_VERIFY] batched MoE shuffle matches the " + "per-expert reference bit-exactly" + ) + + def process_weights_after_loading_mxfp8_moe(self, layer) -> None: """Shuffle weights and scales into FlashInfer TRTLLM MXFP8 layout. @@ -1206,31 +1239,21 @@ def process_weights_after_loading_mxfp8_moe(self, layer) -> None: w13_weight, w2_weight, w13_scale, w2_scale, is_gated, epilogue_tile_m ) - global mxfp8_shuffle_verified - if ( - use_batched_shuffle - and os.getenv("NRL_MXFP8_SHUFFLE_VERIFY") == "1" - and not mxfp8_shuffle_verified - ): - reference = _shuffle_mxfp8_moe_per_expert( - w13_weight, w2_weight, w13_scale, w2_scale, is_gated, epilogue_tile_m - ) - batched = ( - w13_weight_shuffled, - w2_weight_shuffled, - w13_scale_shuffled, - w2_scale_shuffled, - ) - for got, want, tensor_name in zip( - batched, reference, ("w13_weight", "w2_weight", "w13_scale", "w2_scale") - ): - assert torch.equal(got.view(torch.uint8), want.view(torch.uint8)), ( - f"Batched MXFP8 shuffle mismatch vs per-expert reference: {tensor_name}" - ) - mxfp8_shuffle_verified = True - print( - "[NRL_MXFP8_SHUFFLE_VERIFY] batched MoE shuffle matches the " - "per-expert reference bit-exactly" + if use_batched_shuffle: + _verify_mxfp8_moe_shuffle( + layer, + w13_weight, + w2_weight, + w13_scale, + w2_scale, + is_gated, + epilogue_tile_m, + ( + w13_weight_shuffled, + w2_weight_shuffled, + w13_scale_shuffled, + w2_scale_shuffled, + ), ) if first_load: diff --git a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py index e2cea680595..5772c07e281 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py @@ -70,15 +70,14 @@ def mxfp8_e4m3_quantize_for_refit( the torch reference elsewhere. """ x_q = x_scales = None - # Kernel dispatch keys off the TRAINER GPU while the receiver keys off the - # inference GPU. On homogeneous clusters both take the same path; on mixed - # Hopper/Blackwell clusters the flashinfer and torch paths may differ in - # boundary rounding - validate with the parity test before relying on it. if x.is_cuda and torch.cuda.get_device_capability(x.device) >= (10, 0): try: from flashinfer import mxfp8_quantize as flashinfer_mxfp8_quantize - except ImportError: - pass + except ImportError as exc: + raise RuntimeError( + "Trainer-side MXFP8 refit prequantization on sm100+ requires " + "FlashInfer so it matches the vLLM receiver quantization path." + ) from exc else: x_q, x_scales = flashinfer_mxfp8_quantize( x, is_sf_swizzled_layout=False, alignment=32 diff --git a/nemo_rl/models/generation/vllm/vllm_generation.py b/nemo_rl/models/generation/vllm/vllm_generation.py index 9a49fb49e4f..0343ae970ac 100644 --- a/nemo_rl/models/generation/vllm/vllm_generation.py +++ b/nemo_rl/models/generation/vllm/vllm_generation.py @@ -37,7 +37,10 @@ GenerationInterface, GenerationOutputSpec, ) -from nemo_rl.models.generation.vllm.config import VllmConfig +from nemo_rl.models.generation.vllm.config import ( + VllmConfig, + validate_vllm_quantization_config, +) from nemo_rl.models.generation.vllm.utils import ( aggregate_spec_decode_counters, compute_spec_decode_metrics, @@ -100,6 +103,8 @@ def __init__( workers_per_node: Workers per node override defer_model_load: If True, defer model loading for overlapped init """ + validate_vllm_quantization_config(config) + # Store config self.cfg = config self._defer_model_load = defer_model_load diff --git a/nemo_rl/models/policy/interfaces.py b/nemo_rl/models/policy/interfaces.py index e076e986454..6c1c1b5aeaf 100644 --- a/nemo_rl/models/policy/interfaces.py +++ b/nemo_rl/models/policy/interfaces.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. from abc import ABC, abstractmethod -from typing import Any, Optional, TypedDict +from typing import TYPE_CHECKING, Any, Optional, TypedDict import ray import torch @@ -22,6 +22,9 @@ from nemo_rl.models.generation.interfaces import GenerationDatumSpec from nemo_rl.utils.timer import Timer +if TYPE_CHECKING: + from nemo_rl.models.policy import PolicyConfig + class LogprobOutputSpec(TypedDict): """logprobs: Tensor of log probabilities.""" @@ -164,6 +167,8 @@ def shutdown(self) -> bool: class ColocatablePolicyInterface(PolicyInterface): + cfg: "PolicyConfig" + @abstractmethod def init_collective( self, ip: str, port: int, world_size: int, *, train_world_size: int diff --git a/nemo_rl/weight_sync/checkpoint_engine_weight_synchronizer.py b/nemo_rl/weight_sync/checkpoint_engine_weight_synchronizer.py index edf12b0eca6..54419bde3a9 100644 --- a/nemo_rl/weight_sync/checkpoint_engine_weight_synchronizer.py +++ b/nemo_rl/weight_sync/checkpoint_engine_weight_synchronizer.py @@ -20,7 +20,10 @@ from nemo_rl.models.generation.interfaces import CheckpointEngineConfig from nemo_rl.utils.timer import Timer -from nemo_rl.weight_sync.interfaces import WeightSynchronizer +from nemo_rl.weight_sync.interfaces import ( + WeightSynchronizer, + initialize_refit_metadata, +) _MEBIBYTE = 1024 * 1024 @@ -67,22 +70,7 @@ class CheckpointEngineWeightSynchronizer(WeightSynchronizer): _bucket_size_bytes: int | None = None def init_communicator(self) -> None: - state_dict_info = self._policy.prepare_refit_info() - prequant_names = self._generation.prepare_refit_info(state_dict_info) - if prequant_names: - megatron_cfg = self._policy.cfg.get("megatron_cfg", {}) - if not megatron_cfg.get("enabled", False): - raise ValueError( - "vllm_cfg.refit_prequantize requires the Megatron policy backend " - "(policy.megatron_cfg.enabled=true); the DTensor workers do not " - "implement trainer-side pre-quantized refit." - ) - updated_info = self._policy.enable_refit_prequantize(prequant_names) - if updated_info is None: - raise RuntimeError( - "Trainer-side refit prequantization did not return updated metadata." - ) - self._generation.prepare_refit_info(updated_info) + initialize_refit_metadata(self._policy, self._generation) self._ensure_checkpoint_engine_ready() @property diff --git a/nemo_rl/weight_sync/collective_weight_synchronizer.py b/nemo_rl/weight_sync/collective_weight_synchronizer.py index a047915c81c..13f23a1ee17 100644 --- a/nemo_rl/weight_sync/collective_weight_synchronizer.py +++ b/nemo_rl/weight_sync/collective_weight_synchronizer.py @@ -34,7 +34,10 @@ import ray from nemo_rl.utils.timer import Timer -from nemo_rl.weight_sync.interfaces import WeightSynchronizer +from nemo_rl.weight_sync.interfaces import ( + WeightSynchronizer, + initialize_refit_metadata, +) class CollectiveWeightSynchronizer(WeightSynchronizer): @@ -105,8 +108,7 @@ def init_communicator(self) -> None: # prepare_refit_info is called before init_collective. This matches # distillation.py ordering. Neither call depends on the other today, # but we document this as the canonical ordering for future reference. - state_dict_info = self._policy.prepare_refit_info() - self._generation.prepare_refit_info(state_dict_info) + initialize_refit_metadata(self._policy, self._generation) ip, port = self._train_cluster.get_master_address_and_port() train_world_size = self._train_cluster.world_size() diff --git a/nemo_rl/weight_sync/http_weight_synchronizer.py b/nemo_rl/weight_sync/http_weight_synchronizer.py index 13cba3849b7..25f0e6795e5 100644 --- a/nemo_rl/weight_sync/http_weight_synchronizer.py +++ b/nemo_rl/weight_sync/http_weight_synchronizer.py @@ -33,7 +33,10 @@ import ray from nemo_rl.utils.timer import Timer -from nemo_rl.weight_sync.interfaces import WeightSynchronizer +from nemo_rl.weight_sync.interfaces import ( + WeightSynchronizer, + initialize_refit_metadata, +) class HTTPWeightSynchronizer(WeightSynchronizer): @@ -104,8 +107,7 @@ def mark_stale(self) -> None: self._stale = True def init_communicator(self) -> None: - state_dict_info = self._policy.prepare_refit_info() - self._generation.prepare_refit_info(state_dict_info) + initialize_refit_metadata(self._policy, self._generation) def shutdown(self) -> None: pass diff --git a/nemo_rl/weight_sync/interfaces.py b/nemo_rl/weight_sync/interfaces.py index 5b7623e7c68..a26f324f013 100644 --- a/nemo_rl/weight_sync/interfaces.py +++ b/nemo_rl/weight_sync/interfaces.py @@ -38,10 +38,39 @@ """ from abc import ABC, abstractmethod -from typing import Optional +from typing import TYPE_CHECKING, Optional from nemo_rl.utils.timer import Timer +if TYPE_CHECKING: + from nemo_rl.models.generation.interfaces import GenerationInterface + from nemo_rl.models.policy.interfaces import ColocatablePolicyInterface + + +def initialize_refit_metadata( + policy: "ColocatablePolicyInterface", generation: "GenerationInterface" +) -> None: + """Negotiate the wire-format metadata used by policy-to-generation refit.""" + state_dict_info = policy.prepare_refit_info() + prequant_names = generation.prepare_refit_info(state_dict_info) + if not prequant_names: + return + + megatron_cfg = policy.cfg.get("megatron_cfg") + if megatron_cfg is None or not megatron_cfg["enabled"]: + raise ValueError( + "vllm_cfg.refit_prequantize requires the Megatron policy backend " + "(policy.megatron_cfg.enabled=true); the DTensor workers do not " + "implement trainer-side pre-quantized refit." + ) + + updated_info = policy.enable_refit_prequantize(prequant_names) + if updated_info is None: + raise RuntimeError( + "Trainer-side refit prequantization did not return updated metadata." + ) + generation.prepare_refit_info(updated_info) + class WeightSynchronizer(ABC): """Abstract base class for weight synchronization between policy and generation. diff --git a/nemo_rl/weight_sync/ipc_weight_synchronizer.py b/nemo_rl/weight_sync/ipc_weight_synchronizer.py index ab62a9d6782..71ce4e77a82 100644 --- a/nemo_rl/weight_sync/ipc_weight_synchronizer.py +++ b/nemo_rl/weight_sync/ipc_weight_synchronizer.py @@ -34,7 +34,10 @@ import ray from nemo_rl.utils.timer import Timer -from nemo_rl.weight_sync.interfaces import WeightSynchronizer +from nemo_rl.weight_sync.interfaces import ( + WeightSynchronizer, + initialize_refit_metadata, +) class IPCWeightSynchronizer(WeightSynchronizer): @@ -112,8 +115,7 @@ def mark_stale(self) -> None: self._stale = True def init_communicator(self) -> None: - state_dict_info = self._policy.prepare_refit_info() - self._generation.prepare_refit_info(state_dict_info) + initialize_refit_metadata(self._policy, self._generation) def shutdown(self) -> None: pass diff --git a/tests/functional/L1_Functional_Tests_GB200_MXFP8.sh b/tests/functional/L1_Functional_Tests_GB200_MXFP8.sh index 34910911f33..817c0e02361 100644 --- a/tests/functional/L1_Functional_Tests_GB200_MXFP8.sh +++ b/tests/functional/L1_Functional_Tests_GB200_MXFP8.sh @@ -34,6 +34,8 @@ run_test() { fi } +run_test fast uv run --no-sync pytest -q \ + tests/unit/models/generation/test_mxfp8_prequant.py run_test uv run --no-sync bash ./tests/functional/grpo_vllm_mxfp8_rollout_gb200.sh cd ${PROJECT_ROOT}/tests diff --git a/tests/functional/grpo_vllm_mxfp8_rollout_gb200.sh b/tests/functional/grpo_vllm_mxfp8_rollout_gb200.sh index 50639276b70..f459f0c0777 100644 --- a/tests/functional/grpo_vllm_mxfp8_rollout_gb200.sh +++ b/tests/functional/grpo_vllm_mxfp8_rollout_gb200.sh @@ -13,6 +13,7 @@ LOG_DIR=$EXP_DIR/logs JSON_METRICS=$EXP_DIR/metrics.json RUN_LOG=$EXP_DIR/run.log export PYTHONPATH=${PROJECT_ROOT}:${PYTHONPATH:-} +export NRL_MXFP8_SHUFFLE_VERIFY=1 rm -rf "$EXP_DIR" "$LOG_DIR" mkdir -p "$EXP_DIR" "$LOG_DIR" @@ -37,9 +38,12 @@ uv run coverage run -a --data-file="$PROJECT_ROOT/tests/.coverage" --source="$PR policy.train_micro_batch_size=1 \ policy.logprob_batch_size=1 \ policy.max_total_sequence_length=256 \ + policy.dtensor_cfg.enabled=false \ + policy.megatron_cfg.enabled=true \ policy.generation.max_new_tokens=128 \ policy.generation.vllm_cfg.precision=fp8 \ ++policy.generation.vllm_cfg.is_mx=true \ + policy.generation.vllm_cfg.refit_prequantize=true \ policy.generation.vllm_cfg.kv_cache_dtype=auto \ policy.generation.vllm_cfg.max_model_len=256 \ policy.generation.vllm_cfg.gpu_memory_utilization=0.5 \ diff --git a/tests/test_mxfp8_rollout_recipes.py b/tests/test_mxfp8_rollout_recipes.py index 40c04a8d357..187d6ba21f3 100644 --- a/tests/test_mxfp8_rollout_recipes.py +++ b/tests/test_mxfp8_rollout_recipes.py @@ -122,6 +122,8 @@ def test_mxfp8_rollout_recipe_matrix(case_name: str, expected: dict) -> None: assert vllm_cfg["precision"] == "fp8" assert vllm_cfg["is_mx"] is True + assert vllm_cfg["refit_prequantize"] is True + assert config["policy"]["megatron_cfg"]["enabled"] is True assert vllm_cfg["quantization_ignored_layer_kws"] == [ "q_proj", "k_proj", diff --git a/tests/unit/algorithms/test_utils.py b/tests/unit/algorithms/test_utils.py index a22decdf5e0..d4523496488 100755 --- a/tests/unit/algorithms/test_utils.py +++ b/tests/unit/algorithms/test_utils.py @@ -14,7 +14,6 @@ import math from datetime import datetime -from unittest.mock import MagicMock, call import pytest import torch @@ -25,7 +24,6 @@ WALL_CLOCK_EFFICIENCY_CATEGORIES, calculate_baseline_and_std_per_prompt, get_tokenizer, - maybe_enable_refit_prequantize, maybe_pad_last_batch, print_efficiency_summary, print_performance_metrics, @@ -739,64 +737,3 @@ def test_wall_waste_clamped_to_wall_time(self): assert result["efficiency/total_waste_s"] == 60.0 assert result["efficiency/productive_time_s"] == 0.0 assert result["efficiency/efficiency_pct"] == 0.0 - - -class TestMaybeEnableRefitPrequantize: - def test_returns_when_generation_does_not_request_prequantization(self): - policy = MagicMock() - generation = MagicMock() - generation.prepare_refit_info.return_value = None - state_dict_info = {"model.weight": ((2, 2), torch.bfloat16)} - - maybe_enable_refit_prequantize( - policy, - generation, - state_dict_info, - {"megatron_cfg": {"enabled": True}}, - ) - - generation.prepare_refit_info.assert_called_once_with(state_dict_info) - policy.enable_refit_prequantize.assert_not_called() - - @pytest.mark.parametrize( - "policy_config", - [{}, {"megatron_cfg": {"enabled": False}}], - ) - def test_rejects_prequantization_without_megatron(self, policy_config): - policy = MagicMock() - generation = MagicMock() - generation.prepare_refit_info.return_value = ["model.weight"] - - with pytest.raises(ValueError, match="requires the Megatron policy backend"): - maybe_enable_refit_prequantize( - policy, - generation, - {"model.weight": ((2, 2), torch.bfloat16)}, - policy_config, - ) - - policy.enable_refit_prequantize.assert_not_called() - - def test_refreshes_generation_metadata_after_prequantization(self): - policy = MagicMock() - generation = MagicMock() - state_dict_info = {"model.weight": ((2, 2), torch.bfloat16)} - updated_info = { - "model.weight": ((2, 2), torch.float8_e4m3fn), - "model.weight_scale_from_checkpoint": ((2, 1), torch.uint8), - } - generation.prepare_refit_info.side_effect = [["model.weight"], None] - policy.enable_refit_prequantize.return_value = updated_info - - maybe_enable_refit_prequantize( - policy, - generation, - state_dict_info, - {"megatron_cfg": {"enabled": True}}, - ) - - policy.enable_refit_prequantize.assert_called_once_with(["model.weight"]) - assert generation.prepare_refit_info.call_args_list == [ - call(state_dict_info), - call(updated_info), - ] diff --git a/tests/unit/models/generation/test_mxfp8_prequant.py b/tests/unit/models/generation/test_mxfp8_prequant.py index b8a83a4dafb..4d3f848c825 100644 --- a/tests/unit/models/generation/test_mxfp8_prequant.py +++ b/tests/unit/models/generation/test_mxfp8_prequant.py @@ -12,6 +12,8 @@ # See the License for the specific language governing permissions and # limitations under the License. +import sys + import pytest import torch @@ -73,7 +75,25 @@ def test_last_dim_not_divisible_raises(): _mxfp8_e4m3_quantize_torch(x) -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_blackwell_refit_prequantization_requires_flashinfer(monkeypatch): + class FakeBlackwellTensor: + is_cuda = True + device = "cuda" + + monkeypatch.setattr(torch.cuda, "get_device_capability", lambda _device: (10, 0)) + monkeypatch.setitem(sys.modules, "flashinfer", None) + + with pytest.raises(RuntimeError, match=r"sm100\+ requires FlashInfer"): + mxfp8_e4m3_quantize_for_refit(FakeBlackwellTensor()) + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.get_device_capability() < (10, 0), + reason=( + "requires sm100+; below it both sides fall back to the shared torch " + "reference and the comparison is vacuous" + ), +) def test_refit_quantize_matches_receiver_path(): """Bitwise parity with the vLLM receiver path (mxfp8_e4m3_quantize + squeeze).""" vllm_mxfp8 = pytest.importorskip( @@ -82,9 +102,12 @@ def test_refit_quantize_matches_receiver_path(): torch.manual_seed(0) x = torch.randn(256, 512, dtype=torch.bfloat16, device="cuda") + x[0].zero_() ref_lp, ref_scale = vllm_mxfp8.mxfp8_e4m3_quantize(x) ref_scale = torch.squeeze(ref_scale, dim=-1) + assert torch.any(ref_scale == 0) + ref_scale = torch.where(ref_scale == 0, torch.ones_like(ref_scale), ref_scale) got_lp, got_scale = mxfp8_e4m3_quantize_for_refit(x) @@ -149,3 +172,42 @@ def rand_bytes(*shape): assert got.shape == want.shape, name assert got.dtype == want.dtype, name assert torch.equal(got.view(torch.uint8), want.view(torch.uint8)), name + + +def test_mxfp8_shuffle_verification_runs_once_per_layer(monkeypatch): + fp8 = pytest.importorskip("nemo_rl.models.generation.vllm.quantization.fp8") + + class Layer: + pass + + monkeypatch.setenv("NRL_MXFP8_SHUFFLE_VERIFY", "1") + fp8.mxfp8_shuffle_verified_layers.clear() + tensors = ( + torch.arange(8, dtype=torch.uint8), + torch.arange(8, dtype=torch.uint8), + torch.arange(8, dtype=torch.uint8), + torch.arange(8, dtype=torch.uint8), + ) + calls = [] + + def reference(*args): + calls.append(args[0]) + return tensors + + monkeypatch.setattr(fp8, "_shuffle_mxfp8_moe_per_expert", reference) + layers = [Layer(), Layer()] + for layer in layers: + for _ in range(2): + fp8._verify_mxfp8_moe_shuffle( + layer, + tensors[0], + tensors[1], + tensors[2], + tensors[3], + False, + 128, + tensors, + ) + + assert len(calls) == len(layers) + assert all(layer in fp8.mxfp8_shuffle_verified_layers for layer in layers) diff --git a/tests/unit/models/generation/test_vllm_config.py b/tests/unit/models/generation/test_vllm_config.py new file mode 100644 index 00000000000..926e065fcaf --- /dev/null +++ b/tests/unit/models/generation/test_vllm_config.py @@ -0,0 +1,109 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from types import SimpleNamespace +from typing import cast + +import pytest + +from nemo_rl.models.generation import configure_generation_config +from nemo_rl.models.generation.interfaces import GenerationConfig +from nemo_rl.models.generation.vllm.config import ( + VllmConfig, + validate_vllm_quantization_config, +) + + +@pytest.mark.parametrize( + "generation_config", + [ + { + "vllm_cfg": { + "precision": "fp8", + "is_mx": False, + "refit_prequantize": True, + } + }, + { + "vllm_cfg": { + "precision": "fp8", + "refit_prequantize": True, + } + }, + { + "vllm_cfg": { + "precision": "bfloat16", + "refit_prequantize": True, + }, + "quant_cfg": "examples/modelopt/quant_configs/nvfp4_a16.yaml", + "real_quant": True, + }, + ], +) +def test_refit_prequantize_requires_mxfp8(generation_config: dict) -> None: + with pytest.raises( + ValueError, + match="refit_prequantize requires precision='fp8' and is_mx=true", + ): + validate_vllm_quantization_config(cast(VllmConfig, generation_config)) + + +def test_refit_prequantize_must_be_boolean() -> None: + generation_config = cast( + VllmConfig, + { + "vllm_cfg": { + "precision": "fp8", + "is_mx": True, + "refit_prequantize": "false", + } + }, + ) + + with pytest.raises(ValueError, match="refit_prequantize must be a boolean"): + validate_vllm_quantization_config(generation_config) + + +def test_refit_prequantize_accepts_mxfp8() -> None: + generation_config = cast( + VllmConfig, + { + "vllm_cfg": { + "precision": "fp8", + "is_mx": True, + "refit_prequantize": True, + } + }, + ) + + validate_vllm_quantization_config(generation_config) + + +def test_configure_generation_config_validates_refit_prequantize() -> None: + generation_config = cast( + GenerationConfig, + { + "backend": "vllm", + "stop_token_ids": None, + "stop_strings": None, + "vllm_cfg": { + "precision": "bfloat16", + "refit_prequantize": True, + }, + }, + ) + tokenizer = SimpleNamespace(pad_token_id=0, eos_token_id=1) + + with pytest.raises(ValueError, match="requires precision='fp8' and is_mx=true"): + configure_generation_config(generation_config, tokenizer) diff --git a/tests/unit/reference_configs/grpo_math_1B.yaml b/tests/unit/reference_configs/grpo_math_1B.yaml index d4d27703417..f43249f4708 100644 --- a/tests/unit/reference_configs/grpo_math_1B.yaml +++ b/tests/unit/reference_configs/grpo_math_1B.yaml @@ -371,6 +371,8 @@ policy: vllm_cfg: async_engine: false precision: ${policy.precision} + # MXFP8 + Megatron only: quantize on trainer and stream E4M3 values plus scales. + refit_prequantize: false kv_cache_dtype: "auto" tensor_parallel_size: 1 pipeline_parallel_size: 1 diff --git a/tests/unit/weight_sync/test_vllm_remote_sparse_weight_synchronizer.py b/tests/unit/weight_sync/test_vllm_remote_sparse_weight_synchronizer.py index d55acb5a055..1aa4903d501 100644 --- a/tests/unit/weight_sync/test_vllm_remote_sparse_weight_synchronizer.py +++ b/tests/unit/weight_sync/test_vllm_remote_sparse_weight_synchronizer.py @@ -105,6 +105,16 @@ def test_validate_remote_sparse_refit_accepts_supported_scope(): ({"refit_cfg": {"sparse": {"storage": {"s3_bucket": None}}}}, {}), ({"quant_cfg": "fp8"}, {}), ({"vllm_cfg": {"precision": "fp8", "kv_cache_dtype": "auto"}}, {}), + ( + { + "vllm_cfg": { + "precision": "bfloat16", + "kv_cache_dtype": "auto", + "refit_prequantize": True, + } + }, + {}, + ), ( {"vllm_cfg": {"precision": "bfloat16", "kv_cache_dtype": "fp8_e4m3"}}, {}, diff --git a/tests/unit/weight_sync/test_weight_synchronizer.py b/tests/unit/weight_sync/test_weight_synchronizer.py index e799fc56bde..7380e9e3e31 100644 --- a/tests/unit/weight_sync/test_weight_synchronizer.py +++ b/tests/unit/weight_sync/test_weight_synchronizer.py @@ -14,7 +14,7 @@ """Unit tests for the WeightSynchronizer abstraction and its implementations.""" -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock, call, patch import pytest @@ -30,7 +30,10 @@ from nemo_rl.weight_sync.http_weight_synchronizer import ( HTTPWeightSynchronizer, ) -from nemo_rl.weight_sync.interfaces import WeightSynchronizer +from nemo_rl.weight_sync.interfaces import ( + WeightSynchronizer, + initialize_refit_metadata, +) from nemo_rl.weight_sync.ipc_weight_synchronizer import ( IPCWeightSynchronizer, ) @@ -95,6 +98,73 @@ class IncompleteSync(WeightSynchronizer): IncompleteSync() # type: ignore[abstract] +# --------------------------------------------------------------------------- +# Shared refit metadata handshake +# --------------------------------------------------------------------------- + + +class TestInitializeRefitMetadata: + def test_returns_when_generation_does_not_request_prequantization(self): + policy = _mock_policy() + generation = _mock_generation() + state_dict_info = policy.prepare_refit_info.return_value + + initialize_refit_metadata(policy, generation) + + generation.prepare_refit_info.assert_called_once_with(state_dict_info) + policy.enable_refit_prequantize.assert_not_called() + + def test_rejects_prequantization_without_megatron(self): + policy = _mock_policy() + policy.cfg = {"megatron_cfg": {"enabled": False}} + generation = _mock_generation() + generation.prepare_refit_info.return_value = ["layer_0"] + + with pytest.raises(ValueError, match="requires the Megatron policy backend"): + initialize_refit_metadata(policy, generation) + + policy.enable_refit_prequantize.assert_not_called() + + def test_refreshes_generation_metadata_after_prequantization(self): + policy = _mock_policy() + policy.cfg = {"megatron_cfg": {"enabled": True}} + generation = _mock_generation() + state_dict_info = policy.prepare_refit_info.return_value + updated_info = { + "layer_0": { + "shape": [4096, 4096], + "dtype": "float8_e4m3fn", + }, + "layer_0_scale_from_checkpoint": { + "shape": [4096, 128], + "dtype": "uint8", + }, + } + generation.prepare_refit_info.side_effect = [["layer_0"], None] + policy.enable_refit_prequantize.return_value = updated_info + + initialize_refit_metadata(policy, generation) + + policy.enable_refit_prequantize.assert_called_once_with(["layer_0"]) + assert generation.prepare_refit_info.call_args_list == [ + call(state_dict_info), + call(updated_info), + ] + + def test_rejects_missing_prequantized_metadata(self): + policy = _mock_policy() + policy.cfg = {"megatron_cfg": {"enabled": True}} + policy.enable_refit_prequantize.return_value = None + generation = _mock_generation() + generation.prepare_refit_info.return_value = ["layer_0"] + + with pytest.raises( + RuntimeError, + match="did not return updated metadata", + ): + initialize_refit_metadata(policy, generation) + + # --------------------------------------------------------------------------- # IPCWeightSynchronizer # --------------------------------------------------------------------------- @@ -189,6 +259,29 @@ def test_init_communicator(self): policy.prepare_refit_info.assert_called_once() gen.prepare_refit_info.assert_called_once() + def test_init_communicator_completes_prequantization_handshake(self): + policy = _mock_policy() + policy.cfg = {"megatron_cfg": {"enabled": True}} + updated_info = { + "layer_0": { + "shape": [4096, 4096], + "dtype": "float8_e4m3fn", + } + } + policy.enable_refit_prequantize.return_value = updated_info + gen = _mock_generation() + state_dict_info = policy.prepare_refit_info.return_value + gen.prepare_refit_info.side_effect = [["layer_0"], None] + sync = IPCWeightSynchronizer(policy, gen) + + sync.init_communicator() + + policy.enable_refit_prequantize.assert_called_once_with(["layer_0"]) + assert gen.prepare_refit_info.call_args_list == [ + call(state_dict_info), + call(updated_info), + ] + @patch("nemo_rl.weight_sync.ipc_weight_synchronizer.ray") def test_phase_restoration_on_transfer_failure(self, mock_ray): """offload_after_refit and kv_cache prep run even when transfer raises.""" @@ -415,6 +508,37 @@ def test_init_communicator_sets_up_collective(self, mock_ray): "10.0.0.1", 29500, 6, train_world_size=4 ) + @patch("nemo_rl.weight_sync.collective_weight_synchronizer.ray") + def test_init_communicator_prequantizes_before_collective_setup(self, mock_ray): + mock_ray.get.return_value = [True] + policy = _mock_policy() + policy.cfg = {"megatron_cfg": {"enabled": True}} + updated_info = { + "layer_0": { + "shape": [4096, 4096], + "dtype": "float8_e4m3fn", + } + } + policy.enable_refit_prequantize.return_value = updated_info + gen = _mock_generation() + state_dict_info = policy.prepare_refit_info.return_value + gen.prepare_refit_info.side_effect = [["layer_0"], None] + sync = CollectiveWeightSynchronizer( + policy, + gen, + _mock_cluster(world_size=4), + _mock_cluster(world_size=2), + ) + + sync.init_communicator() + + policy.enable_refit_prequantize.assert_called_once_with(["layer_0"]) + assert gen.prepare_refit_info.call_args_list == [ + call(state_dict_info), + call(updated_info), + ] + policy.init_collective.assert_called_once() + # --------------------------------------------------------------------------- # Factory From 58804999c7419f7bf17ade0a9bfe4909446c008c Mon Sep 17 00:00:00 2001 From: seonjinn Date: Thu, 30 Jul 2026 17:15:26 -0700 Subject: [PATCH 14/76] fix(vllm): preserve MXFP8 refit on vLLM 0.25 Signed-off-by: seonjinn --- .../generation/vllm/quantization/fp8.py | 184 ++++++------------ .../generation/test_vllm_fp8_quantization.py | 156 +++++++++------ 2 files changed, 155 insertions(+), 185 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index 13245dcd8f4..b436d6ea878 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -308,6 +308,12 @@ def init_fp8(vllm_cfg, model_name, model_parallel_size): else: fp8_block_quant_kwargs["ignored_layers"].extend(ignored_layers) print("ignored_layers", fp8_block_quant_kwargs["ignored_layers"]) + if global_fp8_config.is_mx: + # vLLM 0.25 also applies ModelOpt MXFP8 to ParallelLMHead, while refit + # sends lm_head in BF16 without a matching block-scale tensor. + fp8_block_quant_kwargs.setdefault("ignored_layers", []) + if "lm_head" not in fp8_block_quant_kwargs["ignored_layers"]: + fp8_block_quant_kwargs["ignored_layers"].append("lm_head") if "ignored_layers" in fp8_block_quant_kwargs: fp8_block_quant_kwargs["ignore"] = fp8_block_quant_kwargs["ignored_layers"] @@ -478,8 +484,8 @@ def load_weights(weights, model_runner): v.to(torch.float), weight_block_size=FP8_BLOCK_QUANT_KWARGS["weight_block_size"], ) - param_scale = torch.squeeze(param_scale, dim=-1) if global_fp8_config.is_mx: + # vLLM 0.25 returns row-major [M, K / 32] E8M0 scales. # All-zero blocks quantize to E8M0 byte 0, which destabilizes the # TRTLLM MXFP8 kernel; clamp to byte 1 (weights are 0 anyway). param_scale = torch.where( @@ -488,6 +494,7 @@ def load_weights(weights, model_runner): weights_quantized.append([k, param_lp]) weights_quantized.append([k + "_scale_from_checkpoint", param_scale]) else: + param_scale = torch.squeeze(param_scale, dim=-1) weights_quantized.append([k, param_lp]) weights_quantized.append([k + "_scale_inv", param_scale]) # Finally load the weights into vllm. Deferred: importing vllm_backend at @@ -1399,33 +1406,40 @@ def process_weights_after_loading_mxfp8_moe(self, layer: RoutedExperts) -> None: layer.w13_weight.copy_(w13_weight_shuffled) layer.w2_weight.copy_(w2_weight_shuffled) + runtime_w13_scale = getattr(layer, "w13_scale_for_apply", layer.w13_weight_scale) + runtime_w2_scale = getattr(layer, "w2_scale_for_apply", layer.w2_weight_scale) + assert self.moe is layer.moe_config -def _get_mxfp8_moe_activation_param( - quant_method: object, - layer: RoutedExperts, - name: str, - local_num_experts: int, - device: torch.device, -) -> torch.Tensor | None: - value = getattr(layer, name, None) - if value is None: - return None - - cache_name = f"_nrl_{name}" - cached = getattr(quant_method, cache_name, None) - if ( - cached is None - or cached.shape != (local_num_experts,) - or cached.device != device - ): - cached = torch.full( - (local_num_experts,), - float(value), - dtype=torch.float32, - device=device, + if self.moe_kernel is None: + from vllm.model_executor.layers.fused_moe.oracle.fp8 import ( + make_fp8_moe_kernel, + make_fp8_moe_quant_config, + ) + + self.moe_quant_config = make_fp8_moe_quant_config( + fp8_backend=self.mxfp8_backend, + w1_scale=runtime_w13_scale, + w2_scale=runtime_w2_scale, + a1_scale=None, + a2_scale=None, + block_shape=self.weight_block_size, + swiglu_limit=getattr(layer, "swiglu_limit", None), + gemm1_alpha=getattr(layer, "swiglu_alpha", None), + gemm1_beta=getattr(layer, "swiglu_beta", None), + layer=layer, + ) + self.moe_kernel = make_fp8_moe_kernel( + moe_quant_config=self.moe_quant_config, + moe_config=self.moe, + fp8_backend=self.mxfp8_backend, + experts_cls=self.experts_cls, + routing_tables=layer._expert_routing_tables(), + layer=layer, ) - setattr(quant_method, cache_name, cached) - return cached + else: + assert self.moe_quant_config is not None + assert self.moe_quant_config.w1_scale is runtime_w13_scale + assert self.moe_quant_config.w2_scale is runtime_w2_scale def apply_monolithic_mxfp8_moe( @@ -1437,59 +1451,13 @@ def apply_monolithic_mxfp8_moe( ) -> torch.Tensor: """Forward for the FlashInfer TRTLLM MXFP8 MoE with hidden-dim padding. - Mirrors vLLM 0.25.1 ModelOptMxFp8FusedMoE.apply_monolithic, with three - changes: reads the *_for_apply tensors built by + Uses vLLM 0.25.1's modular MoE kernel with the *_for_apply tensors built by process_weights_after_loading_mxfp8_moe when padding is active, pads x's - hidden dim to mxfp8_padded_hidden_size before the kernel and narrows the - output back, and allows RELU2_NO_MUL for non-gated MoEs (Nemotron-3-Nano). + hidden dim to mxfp8_padded_hidden_size before the kernel, and narrows the + output back. """ - from flashinfer.fused_moe import ( - ActivationType, - Fp8QuantizationType, - WeightLayout, - ) - from vllm.model_executor.layers.fused_moe.activation import MoEActivation - from vllm.model_executor.layers.fused_moe.config import RoutingMethodType - from vllm.model_executor.layers.fused_moe.oracle.fp8 import Fp8MoeBackend - from vllm.model_executor.layers.quantization.utils.mxfp8_utils import ( - mxfp8_e4m3_quantize, - ) - from vllm.utils.flashinfer import flashinfer_trtllm_fp8_block_scale_moe - - assert self.mxfp8_backend == Fp8MoeBackend.FLASHINFER_TRTLLM - - moe_config = self.moe - parallel_config = moe_config.moe_parallel_config - if parallel_config.enable_eplb: - raise NotImplementedError( - "EPLB is not supported for FlashInfer TRTLLM MXFP8 MoE backend." - ) - - # Map vLLM MoEActivation to FlashInfer ActivationType. - activation_map = { - MoEActivation.SILU: ActivationType.Swiglu, - MoEActivation.RELU2_NO_MUL: ActivationType.Relu2, - } - if moe_config.activation not in activation_map: - raise NotImplementedError( - "FlashInfer TRTLLM MXFP8 MoE supports only " - f"{list(activation_map)}, got {moe_config.activation}." - ) - fi_activation_type = activation_map[moe_config.activation] - - # DeepSeekV3 routing requires float32 logits; others expect bfloat16. - if moe_config.routing_method == RoutingMethodType.DeepSeekV3: - assert router_logits.dtype == torch.float32, ( - "DeepSeekV3 routing requires float32 router_logits, " - f"got {router_logits.dtype}." - ) - else: - router_logits = router_logits.to(torch.bfloat16) - - # Treat 0 as "unset" for compatibility with ungrouped routing configs. - n_group = layer.num_expert_group or None - topk_group = layer.topk_group or None - + assert self.is_monolithic + assert self.moe_kernel is not None unpadded_hidden_size = x.shape[-1] padded_hidden_size = getattr( layer, "mxfp8_padded_hidden_size", unpadded_hidden_size @@ -1499,59 +1467,19 @@ def apply_monolithic_mxfp8_moe( x, (0, padded_hidden_size - unpadded_hidden_size), value=0.0 ) - hidden_states_mxfp8, hidden_states_scale = mxfp8_e4m3_quantize( + output = self.moe_kernel.apply_monolithic( x, - is_sf_swizzled_layout=False, - ) - local_num_experts = moe_config.num_local_experts - - output = flashinfer_trtllm_fp8_block_scale_moe( - routing_logits=router_logits, - routing_bias=layer.e_score_correction_bias, - hidden_states=hidden_states_mxfp8, - hidden_states_scale=hidden_states_scale, - gemm1_weights=getattr(layer, "w13_weight_for_apply", layer.w13_weight), - gemm1_weights_scale=getattr( - layer, "w13_scale_for_apply", layer.w13_weight_scale - ), - gemm2_weights=getattr(layer, "w2_weight_for_apply", layer.w2_weight), - gemm2_weights_scale=getattr(layer, "w2_scale_for_apply", layer.w2_weight_scale), - num_experts=moe_config.num_experts, - top_k=moe_config.experts_per_token, - # Keep Optional semantics: FlashInfer expects None for non-grouped - # routing (e.g. Qwen3 Renormalize), not 0. - n_group=n_group, - topk_group=topk_group, - intermediate_size=moe_config.intermediate_size_per_partition, - local_expert_offset=parallel_config.ep_rank * local_num_experts, - local_num_experts=local_num_experts, + getattr(layer, "w13_weight_for_apply", layer.w13_weight), + getattr(layer, "w2_weight_for_apply", layer.w2_weight), + router_logits, + activation=layer.activation, + global_num_experts=layer.global_num_experts, + expert_map=layer.expert_map, + apply_router_weight_on_input=layer.apply_router_weight_on_input, + num_expert_group=layer.num_expert_group, + topk_group=layer.topk_group, + e_score_correction_bias=layer.e_score_correction_bias, routed_scaling_factor=layer.routed_scaling_factor, - routing_method_type=moe_config.routing_method, - use_shuffled_weight=True, - weight_layout=WeightLayout.MajorK, - fp8_quantization_type=Fp8QuantizationType.MxFp8, - activation_type=fi_activation_type, - gemm1_alpha=_get_mxfp8_moe_activation_param( - self, - layer, - "swiglu_alpha", - local_num_experts, - x.device, - ), - gemm1_beta=_get_mxfp8_moe_activation_param( - self, - layer, - "swiglu_beta", - local_num_experts, - x.device, - ), - gemm1_clamp_limit=_get_mxfp8_moe_activation_param( - self, - layer, - "swiglu_limit", - local_num_experts, - x.device, - ), ) if output.shape[-1] != unpadded_hidden_size: output = output[..., :unpadded_hidden_size].contiguous() diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 3e0bf9c062c..3c8fb63871a 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -74,7 +74,13 @@ def test_init_fp8_uses_mxfp8_quantization_config(fp8_module, monkeypatch): assert vllm_kwargs == { "quantization": "fp8", "kv_cache_dtype": "auto", - "hf_overrides": {"quantization_config": fp8.MXFP8_BLOCK_QUANT_KWARGS}, + "hf_overrides": { + "quantization_config": { + **fp8.MXFP8_BLOCK_QUANT_KWARGS, + "ignored_layers": ["lm_head"], + "ignore": ["lm_head"], + } + }, } assert applied_configs == [fp8.global_fp8_config] assert fp8.global_fp8_config.is_mx is True @@ -206,9 +212,9 @@ def test_load_weights_preserves_prequantized_mxfp8_and_clamps_scales( fp8.global_fp8_config = types.SimpleNamespace(is_mx=True) native = torch.ones(2, 2, dtype=torch.bfloat16) prequantized = torch.ones(2, 2, dtype=torch.float8_e4m3fn) - receiver_quantized = torch.full((2, 2), 2.0, dtype=torch.bfloat16) - receiver_fp8 = torch.ones(2, 2, dtype=torch.float8_e4m3fn) - receiver_scales = torch.tensor([[[0], [7]], [[3], [0]]], dtype=torch.uint8) + receiver_quantized = torch.full((2, 64), 2.0, dtype=torch.bfloat16) + receiver_fp8 = torch.ones(2, 64, dtype=torch.float8_e4m3fn) + receiver_scales = torch.tensor([[0, 7], [3, 0]], dtype=torch.uint8) loaded = [] monkeypatch.setattr( @@ -219,7 +225,7 @@ def test_load_weights_preserves_prequantized_mxfp8_and_clamps_scales( monkeypatch.setattr( mxfp8_utils, "mxfp8_e4m3_quantize", - lambda tensor: ( + lambda tensor, **_kwargs: ( ( receiver_fp8, receiver_scales, @@ -315,10 +321,25 @@ def test_process_mxfp8_moe_pads_kernel_tensors_without_changing_checkpoint_layou fp8_module: types.ModuleType, monkeypatch: pytest.MonkeyPatch, ) -> None: + from vllm.model_executor.layers.fused_moe.oracle import fp8 as fp8_oracle from vllm.model_executor.layers.fused_moe.oracle.fp8 import Fp8MoeBackend fp8 = fp8_module captured: dict[str, Any] = {} + kernel_builds = 0 + + def fake_make_quant_config(**kwargs: Any) -> Any: + captured["quant_config_kwargs"] = kwargs + return types.SimpleNamespace( + w1_scale=kwargs["w1_scale"], + w2_scale=kwargs["w2_scale"], + ) + + def fake_make_kernel(**kwargs: Any) -> Any: + nonlocal kernel_builds + kernel_builds += 1 + captured["kernel_kwargs"] = kwargs + return types.SimpleNamespace() def fake_batched_shuffle( layer: torch.nn.Module, @@ -345,6 +366,8 @@ def fake_batched_shuffle( ) monkeypatch.setattr(fp8, "_shuffle_mxfp8_moe_batched", fake_batched_shuffle) + monkeypatch.setattr(fp8_oracle, "make_fp8_moe_quant_config", fake_make_quant_config) + monkeypatch.setattr(fp8_oracle, "make_fp8_moe_kernel", fake_make_kernel) monkeypatch.delenv("NRL_MXFP8_BATCHED_SHUFFLE", raising=False) layer = torch.nn.Module() @@ -356,6 +379,14 @@ def fake_batched_shuffle( torch.arange(30, dtype=torch.float32).reshape(2, 5, 3), requires_grad=False, ) + layer.w13_weight_scale = torch.nn.Parameter( + torch.zeros(2, 3, 1, dtype=torch.uint8), + requires_grad=False, + ) + layer.w2_weight_scale = torch.nn.Parameter( + torch.zeros(2, 5, 1, dtype=torch.uint8), + requires_grad=False, + ) layer.w13_weight_scale_from_checkpoint = torch.nn.Parameter( torch.zeros(2, 3, 1, dtype=torch.uint8), requires_grad=False, @@ -364,15 +395,19 @@ def fake_batched_shuffle( torch.zeros(2, 5, 1, dtype=torch.uint8), requires_grad=False, ) - layer.moe_config = types.SimpleNamespace(intermediate_size_per_partition=3) + moe_config = types.SimpleNamespace( + intermediate_size_per_partition=3, + is_act_and_mul=False, + ) + layer.moe_config = moe_config + layer._expert_routing_tables = lambda: None quant_method = types.SimpleNamespace( mxfp8_backend=Fp8MoeBackend.FLASHINFER_TRTLLM, experts_cls=types.SimpleNamespace(is_monolithic=lambda: True), weight_block_size=[1, 32], - moe=types.SimpleNamespace( - is_act_and_mul=False, - intermediate_size_per_partition=3, - ), + moe=moe_config, + moe_kernel=None, + moe_quant_config=None, ) original_w13 = layer.w13_weight.detach().clone() original_w2 = layer.w2_weight.detach().clone() @@ -403,6 +438,21 @@ def fake_batched_shuffle( assert layer.w13_scale_for_apply.shape == (2, 128, 16) assert layer.w2_scale_for_apply.shape == (2, 512, 4) assert layer.weight_block_size == [1, 32] + assert captured["quant_config_kwargs"]["w1_scale"] is layer.w13_scale_for_apply + assert captured["quant_config_kwargs"]["w2_scale"] is layer.w2_scale_for_apply + assert captured["kernel_kwargs"]["moe_config"] is moe_config + assert captured["kernel_kwargs"]["routing_tables"] is None + assert kernel_builds == 1 + + w13_scale_for_apply = layer.w13_scale_for_apply + w2_scale_for_apply = layer.w2_scale_for_apply + fp8.process_weights_after_loading_mxfp8_moe(quant_method, layer) + + assert layer.w13_scale_for_apply is w13_scale_for_apply + assert layer.w2_scale_for_apply is w2_scale_for_apply + assert quant_method.moe_quant_config.w1_scale is w13_scale_for_apply + assert quant_method.moe_quant_config.w2_scale is w2_scale_for_apply + assert kernel_builds == 1 def test_process_mxfp8_moe_rejects_non_trtllm_backend_before_mutation( @@ -428,63 +478,50 @@ def test_process_mxfp8_moe_rejects_non_trtllm_backend_before_mutation( def test_apply_monolithic_mxfp8_moe_uses_vllm_025_moe_config( fp8_module: types.ModuleType, - monkeypatch: pytest.MonkeyPatch, ) -> None: from vllm.model_executor.layers.fused_moe.activation import MoEActivation - from vllm.model_executor.layers.fused_moe.config import RoutingMethodType - from vllm.model_executor.layers.fused_moe.oracle.fp8 import Fp8MoeBackend - from vllm.model_executor.layers.quantization.utils import mxfp8_utils - from vllm.utils import flashinfer as vllm_flashinfer fp8 = fp8_module captured: dict[str, Any] = {} - monkeypatch.setattr( - mxfp8_utils, - "mxfp8_e4m3_quantize", - lambda tensor, **_kwargs: ( - tensor.to(torch.float8_e4m3fn), - torch.ones( - tensor.shape[0], - tensor.shape[1] // 32, - dtype=torch.uint8, - ), - ), - ) - - def fake_moe(**kwargs: Any) -> torch.Tensor: + def fake_apply( + x: torch.Tensor, + w13_weight: torch.Tensor, + w2_weight: torch.Tensor, + router_logits: torch.Tensor, + **kwargs: Any, + ) -> torch.Tensor: + captured.update( + { + "x": x, + "w13_weight": w13_weight, + "w2_weight": w2_weight, + "router_logits": router_logits, + } + ) captured.update(kwargs) - return torch.zeros_like(kwargs["hidden_states"], dtype=torch.bfloat16) - - monkeypatch.setattr( - vllm_flashinfer, - "flashinfer_trtllm_fp8_block_scale_moe", - fake_moe, - ) + return torch.zeros_like(x, dtype=torch.bfloat16) - parallel_config = types.SimpleNamespace(enable_eplb=False, ep_rank=3) - moe_config = types.SimpleNamespace( - activation=MoEActivation.RELU2_NO_MUL, - routing_method=RoutingMethodType.Renormalize, - moe_parallel_config=parallel_config, - num_experts=32, - experts_per_token=2, - intermediate_size_per_partition=128, - num_local_experts=4, - ) + kernel = types.SimpleNamespace(apply_monolithic=fake_apply) quant_method = types.SimpleNamespace( - mxfp8_backend=Fp8MoeBackend.FLASHINFER_TRTLLM, - moe=moe_config, + is_monolithic=True, + moe_kernel=kernel, ) + runtime_w13 = torch.empty(4, 128, 512, dtype=torch.float8_e4m3fn) + runtime_w2 = torch.empty(4, 512, 128, dtype=torch.float8_e4m3fn) layer = types.SimpleNamespace( + activation=MoEActivation.RELU2_NO_MUL, + global_num_experts=32, + expert_map=None, + apply_router_weight_on_input=False, num_expert_group=0, topk_group=0, routed_scaling_factor=1.0, e_score_correction_bias=None, w13_weight=torch.empty(4, 128, 512, dtype=torch.float8_e4m3fn), - w13_weight_scale=torch.ones(4, 128, 16, dtype=torch.uint8), w2_weight=torch.empty(4, 512, 128, dtype=torch.float8_e4m3fn), - w2_weight_scale=torch.ones(4, 512, 4, dtype=torch.uint8), + w13_weight_for_apply=runtime_w13, + w2_weight_for_apply=runtime_w2, mxfp8_padded_hidden_size=512, ) x = torch.ones(2, 64, dtype=torch.bfloat16) @@ -497,13 +534,18 @@ def fake_moe(**kwargs: Any) -> torch.Tensor: router_logits, ) - assert captured["num_experts"] == 32 - assert captured["top_k"] == 2 - assert captured["intermediate_size"] == 128 - assert captured["local_expert_offset"] == 12 - assert captured["local_num_experts"] == 4 - assert captured["routing_method_type"] == RoutingMethodType.Renormalize - assert captured["hidden_states"].shape == (2, 512) + assert captured["x"].shape == (2, 512) + assert captured["w13_weight"] is runtime_w13 + assert captured["w2_weight"] is runtime_w2 + assert captured["router_logits"] is router_logits + assert captured["activation"] == MoEActivation.RELU2_NO_MUL + assert captured["global_num_experts"] == 32 + assert captured["expert_map"] is None + assert captured["apply_router_weight_on_input"] is False + assert captured["num_expert_group"] == 0 + assert captured["topk_group"] == 0 + assert captured["e_score_correction_bias"] is None + assert captured["routed_scaling_factor"] == 1.0 assert output.shape == x.shape From 123cc07595d4f6d66380569f6c249c522e3fda0d Mon Sep 17 00:00:00 2001 From: seonjinn Date: Thu, 30 Jul 2026 17:30:48 -0700 Subject: [PATCH 15/76] style(vllm): sort MXFP8 backend import Signed-off-by: seonjinn --- nemo_rl/models/generation/vllm/quantization/fp8.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index b436d6ea878..06dc5c3cbc8 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -1226,13 +1226,13 @@ def process_weights_after_loading_mxfp8_moe(self, layer: RoutedExperts) -> None: keep the original in-place shuffle behavior. """ from vllm.model_executor.layers.fused_moe import FusedMoeWeightScaleSupported + from vllm.model_executor.layers.fused_moe.oracle.fp8 import Fp8MoeBackend from vllm.model_executor.layers.quantization.utils.flashinfer_utils import ( swap_w13_to_w31, ) from vllm.model_executor.layers.quantization.utils.mxfp8_utils import ( MXFP8_BLOCK_SIZE, ) - from vllm.model_executor.layers.fused_moe.oracle.fp8 import Fp8MoeBackend from vllm.model_executor.parameter import ModelWeightParameter from vllm.model_executor.utils import set_weight_attrs From ef3fa08c9eaf657e5370576ba95bf5852575619f Mon Sep 17 00:00:00 2001 From: seonjinn Date: Thu, 30 Jul 2026 20:16:04 -0700 Subject: [PATCH 16/76] fix(vllm): allow partial configs in quant validation Signed-off-by: seonjinn --- nemo_rl/models/generation/vllm/config.py | 4 +++- tests/unit/models/generation/test_vllm_config.py | 6 ++++++ 2 files changed, 9 insertions(+), 1 deletion(-) diff --git a/nemo_rl/models/generation/vllm/config.py b/nemo_rl/models/generation/vllm/config.py index 150daa6fd85..858d0e007f5 100644 --- a/nemo_rl/models/generation/vllm/config.py +++ b/nemo_rl/models/generation/vllm/config.py @@ -160,7 +160,9 @@ class VllmConfig(GenerationConfig): def validate_vllm_quantization_config(config: VllmConfig) -> None: """Reject quantization options that would otherwise be silently ignored.""" - vllm_cfg = config["vllm_cfg"] + vllm_cfg = config.get("vllm_cfg") + if vllm_cfg is None: + return refit_prequantize = vllm_cfg.get("refit_prequantize") if refit_prequantize is not None and not isinstance(refit_prequantize, bool): raise ValueError( diff --git a/tests/unit/models/generation/test_vllm_config.py b/tests/unit/models/generation/test_vllm_config.py index 926e065fcaf..bba7eceaec0 100644 --- a/tests/unit/models/generation/test_vllm_config.py +++ b/tests/unit/models/generation/test_vllm_config.py @@ -90,6 +90,12 @@ def test_refit_prequantize_accepts_mxfp8() -> None: validate_vllm_quantization_config(generation_config) +def test_refit_prequantize_validation_allows_omitted_vllm_cfg() -> None: + generation_config = cast(VllmConfig, {"quant_cfg": None}) + + validate_vllm_quantization_config(generation_config) + + def test_configure_generation_config_validates_refit_prequantize() -> None: generation_config = cast( GenerationConfig, From c4a7b06b1bb583b9571bb5f7937a1df46be1b7f8 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 31 Jul 2026 10:36:48 -0700 Subject: [PATCH 17/76] feat(recipe): add async Qwen3 30B MXFP8 rollout Signed-off-by: seonjinn --- ...-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml | 20 +++++++++++++++++++ 1 file changed, 20 insertions(+) create mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml new file mode 100644 index 00000000000..23304d236df --- /dev/null +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml @@ -0,0 +1,20 @@ +defaults: ./grpo-qwen3-30ba3b-4n8g-async-1off.yaml +checkpointing: + checkpoint_dir: results/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout +policy: + generation: + vllm_cfg: + precision: "fp8" + is_mx: true + refit_prequantize: true + quantization_ignored_layer_kws: + - q_proj + - k_proj + - v_proj + - o_proj + vllm_kwargs: + moe_backend: flashinfer_trtllm +logger: + log_dir: logs/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout + wandb: + name: grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout From 69f7995297790dc7cead57d7d6cea45c0cf2a5be Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 31 Jul 2026 11:14:32 -0700 Subject: [PATCH 18/76] feat(recipe): add async Qwen3 235B MXFP8 rollout Signed-off-by: seonjinn --- ...en3-235b-16n8g-async-1off-mxfp8-rollout.yaml | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) create mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml new file mode 100644 index 00000000000..116b7e6aa38 --- /dev/null +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml @@ -0,0 +1,17 @@ +defaults: ./grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml +checkpointing: + checkpoint_dir: results/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout +policy: + generation: + colocated: + resources: + num_nodes: 8 + gpus_per_node: 8 +logger: + log_dir: logs/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout + wandb: + name: grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout +cluster: + gpus_per_node: 8 + num_nodes: 16 + segment_size: 8 From d53bfdf743b8ffd44ca776fc2ac026b9f8dfd329 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 31 Jul 2026 11:31:22 -0700 Subject: [PATCH 19/76] ci: allowlist async recipe result path Signed-off-by: seonjinn --- .../grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml index 116b7e6aa38..2296cfd580f 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml @@ -1,6 +1,6 @@ defaults: ./grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml checkpointing: - checkpoint_dir: results/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout + checkpoint_dir: results/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout # pragma: allowlist secret policy: generation: colocated: From 2405bf556b79924684f3b4042af53cd5defb0c85 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 31 Jul 2026 12:11:23 -0700 Subject: [PATCH 20/76] fix(vllm): initialize async driver FP8 config Signed-off-by: seonjinn --- .../generation/vllm/quantization/fp8.py | 4 +++ .../generation/test_vllm_fp8_quantization.py | 25 +++++++++++++++++++ 2 files changed, 29 insertions(+) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index 06dc5c3cbc8..548096deae9 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -90,8 +90,12 @@ def my_init(*args, **kwargs): def my_run_engine_core(*args, **kwargs): + global global_fp8_config fp8_cfg = kwargs["vllm_config"].nrl_fp8_cfg del kwargs["vllm_config"].nrl_fp8_cfg + # vLLM 0.25 executes the TP0 driver worker in the EngineCore process. Keep + # the config process-local as well as forwarding it to remote TP workers. + global_fp8_config = fp8_cfg monkey_patch_vllm_ray_executor(fp8_cfg) return original_run_engine_core(*args, **kwargs) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 3c8fb63871a..ea3df28dc7b 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -88,6 +88,31 @@ def test_init_fp8_uses_mxfp8_quantization_config(fp8_module, monkeypatch): assert "VLLM_USE_DEEP_GEMM_E8M0" not in fp8.os.environ +def test_async_engine_core_keeps_fp8_config_for_local_driver_worker( + fp8_module, monkeypatch +): + fp8 = fp8_module + config = fp8.FP8Config( + use_fp8_weights=True, + model_parallel_size=2, + is_mx=True, + ) + vllm_config = types.SimpleNamespace(nrl_fp8_cfg=config) + applied_configs = [] + + monkeypatch.setattr( + fp8, + "monkey_patch_vllm_ray_executor", + lambda fp8_config: applied_configs.append(fp8_config), + ) + monkeypatch.setattr(fp8, "original_run_engine_core", lambda **_kwargs: "done") + + assert fp8.my_run_engine_core(vllm_config=vllm_config) == "done" + assert fp8.global_fp8_config is config + assert applied_configs == [config] + assert not hasattr(vllm_config, "nrl_fp8_cfg") + + @pytest.mark.parametrize( ("field", "error"), [ From 0a11e9456bdb3210cd5c9bc0ce3f1ae297eb5555 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 31 Jul 2026 12:44:57 -0700 Subject: [PATCH 21/76] fix(vllm): propagate FP8 config to refit workers Signed-off-by: seonjinn --- .../generation/vllm/quantization/fp8.py | 22 ++++++++++---- .../models/generation/vllm/vllm_backend.py | 5 +++- nemo_rl/models/generation/vllm/vllm_worker.py | 7 ++++- .../generation/vllm/vllm_worker_async.py | 5 +++- .../models/generation/test_vllm_backend.py | 30 ++++++++++++------- .../generation/test_vllm_fp8_quantization.py | 25 ---------------- 6 files changed, 50 insertions(+), 44 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index 548096deae9..2c41847721a 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -14,7 +14,8 @@ import os import weakref -from dataclasses import dataclass, field +from dataclasses import asdict, dataclass, field +from typing import Any from unittest.mock import patch import ray @@ -74,7 +75,7 @@ class FP8State: # Global FP8 config that can be accessed by patched vLLM functions # initialized by 'init_fp8_cfg()' -global_fp8_config: FP8Config = None +global_fp8_config: FP8Config | None = None # Global FP8 state that holds runtime fp8 objects fp8_state: FP8State = FP8State() @@ -90,16 +91,25 @@ def my_init(*args, **kwargs): def my_run_engine_core(*args, **kwargs): - global global_fp8_config fp8_cfg = kwargs["vllm_config"].nrl_fp8_cfg del kwargs["vllm_config"].nrl_fp8_cfg - # vLLM 0.25 executes the TP0 driver worker in the EngineCore process. Keep - # the config process-local as well as forwarding it to remote TP workers. - global_fp8_config = fp8_cfg monkey_patch_vllm_ray_executor(fp8_cfg) return original_run_engine_core(*args, **kwargs) +def serialize_fp8_config() -> dict[str, Any] | None: + if global_fp8_config is None: + return None + return asdict(global_fp8_config) + + +def install_fp8_config(config: dict[str, Any] | None) -> None: + if config is None: + return + global global_fp8_config + global_fp8_config = FP8Config(**config) + + def monkey_patch_vllm_ray_executor(fp8_config): if fp8_config.model_parallel_size > 1: # we patch vllm's collective_rpc so that before vllm initalizes the model on each rank, we execute diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index 8484713ab9a..eed389a936e 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -448,7 +448,9 @@ def maybe_init_zmq(self): self.zmq_socket.connect(self.get_zmq_address()) def prepare_refit_info( - self, state_dict_info: dict[str, Any] + self, + state_dict_info: dict[str, Any], + serialized_fp8_config: Optional[dict[str, Any]] = None, ) -> Optional[list[str]]: """Prepare state dict metadata for weight refitting and IPC streaming. @@ -467,6 +469,7 @@ def prepare_refit_info( from nemo_rl.models.generation.vllm.quantization import fp8 + fp8.install_fp8_config(serialized_fp8_config) if not ( fp8.global_fp8_config is not None and fp8.global_fp8_config.is_mx diff --git a/nemo_rl/models/generation/vllm/vllm_worker.py b/nemo_rl/models/generation/vllm/vllm_worker.py index 1eb29029306..5186bea6811 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker.py +++ b/nemo_rl/models/generation/vllm/vllm_worker.py @@ -1096,7 +1096,12 @@ def prepare_refit_info( Returns the parameter names the engine wants pre-quantized on the trainer (vllm_cfg.refit_prequantize), or None. """ - results = self.llm.collective_rpc("prepare_refit_info", args=(state_dict_info,)) + from nemo_rl.models.generation.vllm.quantization import fp8 + + results = self.llm.collective_rpc( + "prepare_refit_info", + args=(state_dict_info, fp8.serialize_fp8_config()), + ) # Union across the engine's TP/PP workers: with pipeline parallelism # each shard only classifies its local parameters as fp8-eligible. names = sorted({name for result in results if result for name in result}) diff --git a/nemo_rl/models/generation/vllm/vllm_worker_async.py b/nemo_rl/models/generation/vllm/vllm_worker_async.py index aa82c1b638f..86c3ad7eac0 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker_async.py +++ b/nemo_rl/models/generation/vllm/vllm_worker_async.py @@ -1304,8 +1304,11 @@ async def prepare_refit_info_async( self, state_dict_info: dict[str, Any] ) -> Optional[list[str]]: """Async version of prepare_refit_info.""" + from nemo_rl.models.generation.vllm.quantization import fp8 + results = await self.llm.collective_rpc( - "prepare_refit_info", args=(state_dict_info,) + "prepare_refit_info", + args=(state_dict_info, fp8.serialize_fp8_config()), ) # Union across the engine's TP/PP workers: with pipeline parallelism # each shard only classifies its local parameters as fp8-eligible. diff --git a/tests/unit/models/generation/test_vllm_backend.py b/tests/unit/models/generation/test_vllm_backend.py index cb8b5a84f14..05c1a5eecd2 100644 --- a/tests/unit/models/generation/test_vllm_backend.py +++ b/tests/unit/models/generation/test_vllm_backend.py @@ -143,17 +143,21 @@ def is_fp8_weight(name, candidate_model): checked_names.append(name) return name == "model.linear.weight" - monkeypatch.setattr( - fp8, - "global_fp8_config", - SimpleNamespace(is_mx=True, refit_prequantize=enabled), - ) + source_config = fp8.FP8Config( + is_mx=True, + refit_prequantize=enabled, + use_fp8_weights=True, + ) + monkeypatch.setattr(fp8, "global_fp8_config", source_config) + serialized_config = fp8.serialize_fp8_config() + monkeypatch.setattr(fp8, "global_fp8_config", None) monkeypatch.setattr(fp8, "is_fp8_model", is_fp8_model) monkeypatch.setattr(fp8, "_is_fp8_weight", is_fp8_weight) - result = ext.prepare_refit_info(state_dict_info) + result = ext.prepare_refit_info(state_dict_info, serialized_config) assert ext.state_dict_info is state_dict_info + assert fp8.global_fp8_config == source_config if enabled: assert result == ["model.linear.weight"] is_fp8_model.assert_called_once_with(config) @@ -165,13 +169,16 @@ def is_fp8_weight(name, candidate_model): @pytest.mark.vllm -def test_sync_prepare_refit_info_unions_worker_names(): +def test_sync_prepare_refit_info_unions_worker_names(monkeypatch): + from nemo_rl.models.generation.vllm.quantization import fp8 from nemo_rl.models.generation.vllm.vllm_worker import ( VllmGenerationWorkerImpl, ) worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) state_dict_info = {"model.weight": ((2, 2), torch.bfloat16)} + serialized_config = {"is_mx": True, "refit_prequantize": True} + monkeypatch.setattr(fp8, "serialize_fp8_config", lambda: serialized_config) worker.llm = SimpleNamespace( collective_rpc=MagicMock( return_value=[ @@ -188,19 +195,22 @@ def test_sync_prepare_refit_info_unions_worker_names(): ] worker.llm.collective_rpc.assert_called_once_with( "prepare_refit_info", - args=(state_dict_info,), + args=(state_dict_info, serialized_config), ) @pytest.mark.vllm @pytest.mark.asyncio -async def test_async_prepare_refit_info_unions_worker_names(): +async def test_async_prepare_refit_info_unions_worker_names(monkeypatch): + from nemo_rl.models.generation.vllm.quantization import fp8 from nemo_rl.models.generation.vllm.vllm_worker_async import ( VllmAsyncGenerationWorkerImpl, ) worker = VllmAsyncGenerationWorkerImpl.__new__(VllmAsyncGenerationWorkerImpl) state_dict_info = {"model.weight": ((2, 2), torch.bfloat16)} + serialized_config = {"is_mx": True, "refit_prequantize": True} + monkeypatch.setattr(fp8, "serialize_fp8_config", lambda: serialized_config) worker.llm = SimpleNamespace( collective_rpc=AsyncMock( return_value=[None, ["model.b.weight"], ["model.a.weight"]] @@ -213,7 +223,7 @@ async def test_async_prepare_refit_info_unions_worker_names(): ] worker.llm.collective_rpc.assert_awaited_once_with( "prepare_refit_info", - args=(state_dict_info,), + args=(state_dict_info, serialized_config), ) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index ea3df28dc7b..3c8fb63871a 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -88,31 +88,6 @@ def test_init_fp8_uses_mxfp8_quantization_config(fp8_module, monkeypatch): assert "VLLM_USE_DEEP_GEMM_E8M0" not in fp8.os.environ -def test_async_engine_core_keeps_fp8_config_for_local_driver_worker( - fp8_module, monkeypatch -): - fp8 = fp8_module - config = fp8.FP8Config( - use_fp8_weights=True, - model_parallel_size=2, - is_mx=True, - ) - vllm_config = types.SimpleNamespace(nrl_fp8_cfg=config) - applied_configs = [] - - monkeypatch.setattr( - fp8, - "monkey_patch_vllm_ray_executor", - lambda fp8_config: applied_configs.append(fp8_config), - ) - monkeypatch.setattr(fp8, "original_run_engine_core", lambda **_kwargs: "done") - - assert fp8.my_run_engine_core(vllm_config=vllm_config) == "done" - assert fp8.global_fp8_config is config - assert applied_configs == [config] - assert not hasattr(vllm_config, "nrl_fp8_cfg") - - @pytest.mark.parametrize( ("field", "error"), [ From 81ce860a184409cf7a3acbefa27920196a4af398 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 31 Jul 2026 13:31:40 -0700 Subject: [PATCH 22/76] fix(recipe): keep Qwen router gate in BF16 Signed-off-by: seonjinn --- ...n3-235b-16n8g-async-1off-mxfp8-rollout.yaml | 7 +++++++ ...3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml | 1 + tests/test_mxfp8_rollout_recipes.py | 18 ++++++++++++++++++ 3 files changed, 26 insertions(+) diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml index 2296cfd580f..21e54e91224 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml @@ -7,6 +7,13 @@ policy: resources: num_nodes: 8 gpus_per_node: 8 + vllm_cfg: + quantization_ignored_layer_kws: + - q_proj + - k_proj + - v_proj + - o_proj + - ".mlp.gate" logger: log_dir: logs/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout wandb: diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml index 23304d236df..4160e9b478a 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml @@ -12,6 +12,7 @@ policy: - k_proj - v_proj - o_proj + - ".mlp.gate" vllm_kwargs: moe_backend: flashinfer_trtllm logger: diff --git a/tests/test_mxfp8_rollout_recipes.py b/tests/test_mxfp8_rollout_recipes.py index 187d6ba21f3..392172e4fac 100644 --- a/tests/test_mxfp8_rollout_recipes.py +++ b/tests/test_mxfp8_rollout_recipes.py @@ -196,3 +196,21 @@ def test_qwen3_235b_mxfp8_recipes_keep_baseline_runtime_knobs() -> None: assert "max_num_steps" not in grpo_config assert "val_batch_size" not in grpo_config assert "max_val_samples" not in grpo_config + + +@pytest.mark.parametrize( + "case_name", + ( + "grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout", + "grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout", + ), +) +def test_b200_async_mxfp8_recipes_keep_router_gate_in_bf16(case_name: str) -> None: + config = _load_resolved_yaml(PERF_CONFIG_DIR / f"{case_name}.yaml") + ignored = config["policy"]["generation"]["vllm_cfg"][ + "quantization_ignored_layer_kws" + ] + + assert ignored == ["q_proj", "k_proj", "v_proj", "o_proj", ".mlp.gate"] + assert "gate" not in ignored + assert "gate_proj" not in ignored From 9e4337850523b9aefd2acda3f96fa0ca11698749 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 31 Jul 2026 13:41:21 -0700 Subject: [PATCH 23/76] fix(recipe): exclude Qwen MoE routers from MXFP8 Signed-off-by: seonjinn --- .../grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml | 1 + ...3-235b-16n8g-async-1off-mxfp8-rollout.yaml | 7 ------- ...3-235b-32n4g-async-1off-mxfp8-rollout.yaml | 1 + ...-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml | 5 ----- .../grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml | 1 + tests/test_mxfp8_rollout_recipes.py | 19 ++++++++++++++++--- 6 files changed, 19 insertions(+), 15 deletions(-) diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml index 53c313db9a6..ecc303c00fb 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml @@ -13,6 +13,7 @@ policy: - k_proj - v_proj - o_proj + - ".mlp.gate" vllm_kwargs: moe_backend: flashinfer_trtllm logger: diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml index 21e54e91224..2296cfd580f 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml @@ -7,13 +7,6 @@ policy: resources: num_nodes: 8 gpus_per_node: 8 - vllm_cfg: - quantization_ignored_layer_kws: - - q_proj - - k_proj - - v_proj - - o_proj - - ".mlp.gate" logger: log_dir: logs/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout wandb: diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml index 303e543f1e8..4f902936356 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml @@ -13,6 +13,7 @@ policy: - k_proj - v_proj - o_proj + - ".mlp.gate" vllm_kwargs: moe_backend: flashinfer_trtllm logger: diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml index 8d1713d7f15..a3477ca4f8e 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml @@ -22,11 +22,6 @@ policy: precision: "fp8" is_mx: true refit_prequantize: true - quantization_ignored_layer_kws: - - q_proj - - k_proj - - v_proj - - o_proj logger: log_dir: logs/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout wandb: diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml index dc4eafefe9b..6e439bcf907 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml @@ -46,6 +46,7 @@ policy: - k_proj - v_proj - o_proj + - ".mlp.gate" logger: log_dir: logs/grpo-qwen3-30ba3b-4n4g-mxfp8-rollout wandb_enabled: true diff --git a/tests/test_mxfp8_rollout_recipes.py b/tests/test_mxfp8_rollout_recipes.py index 392172e4fac..65a9215df24 100644 --- a/tests/test_mxfp8_rollout_recipes.py +++ b/tests/test_mxfp8_rollout_recipes.py @@ -29,6 +29,7 @@ "segment_size": 4, "async_engine": None, "moe_backend": None, + "router_gate_bf16": True, }, "grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout": { "nodes": 4, @@ -36,6 +37,7 @@ "segment_size": 2, "async_engine": True, "moe_backend": None, + "router_gate_bf16": True, }, "grpo-qwen3-32b-4n4g-mxfp8-rollout": { "nodes": 4, @@ -43,6 +45,7 @@ "segment_size": 4, "async_engine": None, "moe_backend": None, + "router_gate_bf16": False, }, "grpo-qwen3-32b-8n4g-async-1off-mxfp8-rollout": { "nodes": 8, @@ -50,6 +53,7 @@ "segment_size": 4, "async_engine": True, "moe_backend": None, + "router_gate_bf16": False, }, "grpo-qwen3-235b-16n4g-mxfp8-rollout": { "nodes": 16, @@ -58,6 +62,7 @@ "async_engine": True, "tensor_parallel_size": 4, "moe_backend": "flashinfer_trtllm", + "router_gate_bf16": True, }, "grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout": { "nodes": 32, @@ -66,6 +71,7 @@ "async_engine": True, "tensor_parallel_size": 4, "moe_backend": "flashinfer_trtllm", + "router_gate_bf16": True, }, } @@ -124,12 +130,15 @@ def test_mxfp8_rollout_recipe_matrix(case_name: str, expected: dict) -> None: assert vllm_cfg["is_mx"] is True assert vllm_cfg["refit_prequantize"] is True assert config["policy"]["megatron_cfg"]["enabled"] is True - assert vllm_cfg["quantization_ignored_layer_kws"] == [ + expected_ignored = [ "q_proj", "k_proj", "v_proj", "o_proj", ] + if expected["router_gate_bf16"]: + expected_ignored.append(".mlp.gate") + assert vllm_cfg["quantization_ignored_layer_kws"] == expected_ignored assert cluster["num_nodes"] == expected["nodes"] assert cluster["gpus_per_node"] == expected["gpus_per_node"] assert cluster["segment_size"] == expected["segment_size"] @@ -212,5 +221,9 @@ def test_b200_async_mxfp8_recipes_keep_router_gate_in_bf16(case_name: str) -> No ] assert ignored == ["q_proj", "k_proj", "v_proj", "o_proj", ".mlp.gate"] - assert "gate" not in ignored - assert "gate_proj" not in ignored + router_name = "model.layers.0.mlp.gate" + expert_name = "model.layers.0.mlp.experts.0.gate_proj" + fused_expert_name = "model.layers.0.mlp.experts.gate_up_proj" + assert any(keyword in router_name for keyword in ignored) + assert not any(keyword in expert_name for keyword in ignored) + assert not any(keyword in fused_expert_name for keyword in ignored) From 6e582b86f90b3138004882e2b3305e3aad8d1fe4 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 31 Jul 2026 20:33:34 -0700 Subject: [PATCH 24/76] fix(vllm): patch FP8 in RayExecutorV2 workers Signed-off-by: seonjinn --- .../generation/vllm/quantization/fp8.py | 37 +++++++++++++ .../generation/test_vllm_fp8_quantization.py | 55 +++++++++++++++++++ 2 files changed, 92 insertions(+) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index 2c41847721a..1e2a37ef656 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -110,7 +110,44 @@ def install_fp8_config(config: dict[str, Any] | None) -> None: global_fp8_config = FP8Config(**config) +def _patch_ray_executor_v2_worker( + ray_worker_proc_cls: type[Any], fp8_config: FP8Config +) -> None: + """Install FP8 patches inside RayExecutorV2 workers before model loading.""" + original_initialize_worker = ray_worker_proc_cls.initialize_worker + if getattr(original_initialize_worker, "_nrl_fp8_patched", False): + return + + def initialize_worker_with_fp8_patches( + self: Any, + local_rank: int, + env_vars: dict[str, str], + driver_env_vars: dict[str, str] | None = None, + assigned_physical_gpu_ids: list[int] | None = None, + ) -> Any: + global fp8_patches_applied + if not fp8_patches_applied: + apply_fp8_patches(None, fp8_config) + return original_initialize_worker( + self, + local_rank, + env_vars, + driver_env_vars, + assigned_physical_gpu_ids, + ) + + setattr(initialize_worker_with_fp8_patches, "_nrl_fp8_patched", True) + ray_worker_proc_cls.initialize_worker = initialize_worker_with_fp8_patches + + def monkey_patch_vllm_ray_executor(fp8_config): + try: + from vllm.v1.executor.ray_executor_v2 import RayWorkerProc + except ImportError: + pass + else: + _patch_ray_executor_v2_worker(RayWorkerProc, fp8_config) + if fp8_config.model_parallel_size > 1: # we patch vllm's collective_rpc so that before vllm initalizes the model on each rank, we execute # a ray remote that patches each worker with the required fp8 vllm patches diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 3c8fb63871a..1649e03c565 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -88,6 +88,61 @@ def test_init_fp8_uses_mxfp8_quantization_config(fp8_module, monkeypatch): assert "VLLM_USE_DEEP_GEMM_E8M0" not in fp8.os.environ +def test_ray_executor_v2_worker_applies_fp8_patches_before_model_load( + fp8_module, monkeypatch +): + fp8 = fp8_module + events = [] + config = fp8.FP8Config( + use_fp8_weights=True, + model_parallel_size=2, + is_mx=True, + ) + + class FakeRayWorkerProc: + def initialize_worker( + self, + local_rank, + env_vars, + driver_env_vars=None, + assigned_physical_gpu_ids=None, + ): + events.append( + ( + "model_load", + local_rank, + env_vars, + driver_env_vars, + assigned_physical_gpu_ids, + ) + ) + + def fake_apply_fp8_patches(_self, fp8_config): + events.append(("fp8_patches", fp8_config)) + fp8.fp8_patches_applied = True + + monkeypatch.setattr(fp8, "apply_fp8_patches", fake_apply_fp8_patches) + fp8._patch_ray_executor_v2_worker(FakeRayWorkerProc, config) + + FakeRayWorkerProc().initialize_worker( + 1, + {"WORKER_ENV": "1"}, + {"DRIVER_ENV": "1"}, + assigned_physical_gpu_ids=[2, 3], + ) + + assert events == [ + ("fp8_patches", config), + ( + "model_load", + 1, + {"WORKER_ENV": "1"}, + {"DRIVER_ENV": "1"}, + [2, 3], + ), + ] + + @pytest.mark.parametrize( ("field", "error"), [ From 93394c2a2a2403ef8e2093b379bcb7b340a6bc26 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 31 Jul 2026 20:42:03 -0700 Subject: [PATCH 25/76] fix(vllm): serialize FP8 worker pre-init hook Signed-off-by: seonjinn --- .../generation/vllm/quantization/fp8.py | 49 ++++++++++--------- .../generation/test_vllm_fp8_quantization.py | 40 +++++++-------- 2 files changed, 46 insertions(+), 43 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index 1e2a37ef656..dd6b3d7d43c 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -111,42 +111,45 @@ def install_fp8_config(config: dict[str, Any] | None) -> None: def _patch_ray_executor_v2_worker( - ray_worker_proc_cls: type[Any], fp8_config: FP8Config + ray_executor_v2: Any, fp8_config: FP8Config ) -> None: """Install FP8 patches inside RayExecutorV2 workers before model loading.""" - original_initialize_worker = ray_worker_proc_cls.initialize_worker - if getattr(original_initialize_worker, "_nrl_fp8_patched", False): + original_ray_worker_proc = ray_executor_v2.RayWorkerProc + if getattr(original_ray_worker_proc, "_nrl_fp8_patched", False): + original_ray_worker_proc._nrl_fp8_config = fp8_config return - def initialize_worker_with_fp8_patches( - self: Any, - local_rank: int, - env_vars: dict[str, str], - driver_env_vars: dict[str, str] | None = None, - assigned_physical_gpu_ids: list[int] | None = None, - ) -> Any: - global fp8_patches_applied - if not fp8_patches_applied: - apply_fp8_patches(None, fp8_config) - return original_initialize_worker( + class NRLFP8RayWorkerProc(original_ray_worker_proc): + _nrl_fp8_patched = True + _nrl_fp8_config = fp8_config + + def initialize_worker( self, - local_rank, - env_vars, - driver_env_vars, - assigned_physical_gpu_ids, - ) + local_rank: int, + env_vars: dict[str, str], + driver_env_vars: dict[str, str] | None = None, + assigned_physical_gpu_ids: list[int] | None = None, + ) -> Any: + global fp8_patches_applied + if not fp8_patches_applied: + apply_fp8_patches(None, type(self)._nrl_fp8_config) + return super().initialize_worker( + local_rank, + env_vars, + driver_env_vars, + assigned_physical_gpu_ids, + ) - setattr(initialize_worker_with_fp8_patches, "_nrl_fp8_patched", True) - ray_worker_proc_cls.initialize_worker = initialize_worker_with_fp8_patches + ray_executor_v2.RayWorkerProc = NRLFP8RayWorkerProc def monkey_patch_vllm_ray_executor(fp8_config): try: - from vllm.v1.executor.ray_executor_v2 import RayWorkerProc + from vllm.v1.executor import ray_executor_v2 except ImportError: pass else: - _patch_ray_executor_v2_worker(RayWorkerProc, fp8_config) + _patch_ray_executor_v2_worker(ray_executor_v2, fp8_config) if fp8_config.model_parallel_size > 1: # we patch vllm's collective_rpc so that before vllm initalizes the model on each rank, we execute diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 1649e03c565..45d7f310d07 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import cloudpickle import types from typing import Any @@ -107,14 +108,12 @@ def initialize_worker( driver_env_vars=None, assigned_physical_gpu_ids=None, ): - events.append( - ( - "model_load", - local_rank, - env_vars, - driver_env_vars, - assigned_physical_gpu_ids, - ) + assert fp8.fp8_patches_applied + return ( + local_rank, + env_vars, + driver_env_vars, + assigned_physical_gpu_ids, ) def fake_apply_fp8_patches(_self, fp8_config): @@ -122,25 +121,26 @@ def fake_apply_fp8_patches(_self, fp8_config): fp8.fp8_patches_applied = True monkeypatch.setattr(fp8, "apply_fp8_patches", fake_apply_fp8_patches) - fp8._patch_ray_executor_v2_worker(FakeRayWorkerProc, config) + ray_executor_v2 = types.SimpleNamespace(RayWorkerProc=FakeRayWorkerProc) + fp8._patch_ray_executor_v2_worker(ray_executor_v2, config) + patched_worker_cls = cloudpickle.loads( + cloudpickle.dumps(ray_executor_v2.RayWorkerProc) + ) - FakeRayWorkerProc().initialize_worker( + result = patched_worker_cls().initialize_worker( 1, {"WORKER_ENV": "1"}, {"DRIVER_ENV": "1"}, assigned_physical_gpu_ids=[2, 3], ) - assert events == [ - ("fp8_patches", config), - ( - "model_load", - 1, - {"WORKER_ENV": "1"}, - {"DRIVER_ENV": "1"}, - [2, 3], - ), - ] + assert events == [("fp8_patches", config)] + assert result == ( + 1, + {"WORKER_ENV": "1"}, + {"DRIVER_ENV": "1"}, + [2, 3], + ) @pytest.mark.parametrize( From f3c9196a38cdbba7ba2e0f45da8863f1f4711f5d Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 31 Jul 2026 21:29:32 -0700 Subject: [PATCH 26/76] test(vllm): share RayExecutorV2 patch recorder Signed-off-by: seonjinn --- tests/unit/models/generation/test_vllm_fp8_quantization.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 45d7f310d07..94a31564fbd 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -93,7 +93,7 @@ def test_ray_executor_v2_worker_applies_fp8_patches_before_model_load( fp8_module, monkeypatch ): fp8 = fp8_module - events = [] + monkeypatch.setattr(fp8, "_test_applied_configs", [], raising=False) config = fp8.FP8Config( use_fp8_weights=True, model_parallel_size=2, @@ -117,7 +117,7 @@ def initialize_worker( ) def fake_apply_fp8_patches(_self, fp8_config): - events.append(("fp8_patches", fp8_config)) + fp8._test_applied_configs.append(fp8_config) fp8.fp8_patches_applied = True monkeypatch.setattr(fp8, "apply_fp8_patches", fake_apply_fp8_patches) @@ -134,7 +134,7 @@ def fake_apply_fp8_patches(_self, fp8_config): assigned_physical_gpu_ids=[2, 3], ) - assert events == [("fp8_patches", config)] + assert fp8._test_applied_configs == [config] assert result == ( 1, {"WORKER_ENV": "1"}, From e3384d2795f3862e3fd4455d3f4fee96cff72c6a Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sat, 1 Aug 2026 11:41:26 -0700 Subject: [PATCH 27/76] feat(vllm): configure refit runtime optimizations Signed-off-by: seonjinn --- examples/configs/distillation_math.yaml | 2 ++ examples/configs/grpo_math_1B.yaml | 2 ++ examples/configs/ppo_math_1B.yaml | 2 ++ .../grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml | 2 ++ ...3-235b-32n4g-async-1off-mxfp8-rollout.yaml | 2 ++ ...-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml | 2 ++ .../grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml | 2 ++ .../grpo-qwen3-32b-4n4g-mxfp8-rollout.yaml | 2 ++ ...en3-32b-8n4g-async-1off-mxfp8-rollout.yaml | 2 ++ nemo_rl/models/generation/vllm/config.py | 8 +++++++ .../generation/vllm/quantization/fp8.py | 14 ++++++++--- .../models/generation/vllm/vllm_backend.py | 24 ++++++++++++------- tests/test_mxfp8_rollout_recipes.py | 2 ++ .../models/generation/test_vllm_config.py | 20 ++++++++++++++++ .../generation/test_vllm_fp8_quantization.py | 8 ++++++- .../generation/test_vllm_refit_loader.py | 12 ++++++---- .../reference_configs/distillation_math.yaml | 2 ++ .../unit/reference_configs/grpo_math_1B.yaml | 2 ++ .../ppo_math_1B_megatron.yaml | 2 ++ .../weight_sync/test_weight_synchronizer.py | 24 +++++++++++++++++++ 20 files changed, 120 insertions(+), 16 deletions(-) diff --git a/examples/configs/distillation_math.yaml b/examples/configs/distillation_math.yaml index 11ca9d09072..b8bced51dc0 100644 --- a/examples/configs/distillation_math.yaml +++ b/examples/configs/distillation_math.yaml @@ -205,6 +205,8 @@ policy: &POLICY_BASE vllm_cfg: async_engine: false precision: ${...precision} + refit_batched_moe_shuffle: true # Batch MoE layout transforms across experts during refit. + refit_cache_loader_routes: false # Replay stable vLLM loader routes across refits. kv_cache_dtype: "auto" tensor_parallel_size: 1 pipeline_parallel_size: 1 diff --git a/examples/configs/grpo_math_1B.yaml b/examples/configs/grpo_math_1B.yaml index ee464e2fec2..bfdf2fea083 100644 --- a/examples/configs/grpo_math_1B.yaml +++ b/examples/configs/grpo_math_1B.yaml @@ -388,6 +388,8 @@ policy: precision: ${policy.precision} # MXFP8 + Megatron only: quantize on trainer and stream E4M3 values plus scales. refit_prequantize: false + refit_batched_moe_shuffle: true # Batch MoE layout transforms across experts during refit. + refit_cache_loader_routes: false # Replay stable vLLM loader routes across refits. kv_cache_dtype: "auto" tensor_parallel_size: 1 pipeline_parallel_size: 1 diff --git a/examples/configs/ppo_math_1B.yaml b/examples/configs/ppo_math_1B.yaml index ceed0e3c4ef..2dc9a346b08 100644 --- a/examples/configs/ppo_math_1B.yaml +++ b/examples/configs/ppo_math_1B.yaml @@ -247,6 +247,8 @@ policy: vllm_cfg: async_engine: false precision: ${policy.precision} + refit_batched_moe_shuffle: true # Batch MoE layout transforms across experts during refit. + refit_cache_loader_routes: false # Replay stable vLLM loader routes across refits. kv_cache_dtype: "auto" tensor_parallel_size: 1 pipeline_parallel_size: 1 diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml index ecc303c00fb..4af61c4906a 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n4g-mxfp8-rollout.yaml @@ -8,6 +8,8 @@ policy: precision: "fp8" is_mx: true refit_prequantize: true + refit_batched_moe_shuffle: true + refit_cache_loader_routes: true quantization_ignored_layer_kws: - q_proj - k_proj diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml index 4f902936356..78114e0d7ff 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml @@ -8,6 +8,8 @@ policy: precision: "fp8" is_mx: true refit_prequantize: true + refit_batched_moe_shuffle: true + refit_cache_loader_routes: true quantization_ignored_layer_kws: - q_proj - k_proj diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml index a3477ca4f8e..9ea0a7728e4 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml @@ -22,6 +22,8 @@ policy: precision: "fp8" is_mx: true refit_prequantize: true + refit_batched_moe_shuffle: true + refit_cache_loader_routes: true logger: log_dir: logs/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout wandb: diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml index 6e439bcf907..39963484c1c 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-mxfp8-rollout.yaml @@ -41,6 +41,8 @@ policy: precision: "fp8" is_mx: true refit_prequantize: true + refit_batched_moe_shuffle: true + refit_cache_loader_routes: true quantization_ignored_layer_kws: - q_proj - k_proj diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-mxfp8-rollout.yaml index 3107e8ff66f..fb3b7f5958c 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-4n4g-mxfp8-rollout.yaml @@ -9,6 +9,8 @@ policy: precision: "fp8" is_mx: true refit_prequantize: true + refit_batched_moe_shuffle: true + refit_cache_loader_routes: true quantization_ignored_layer_kws: - q_proj - k_proj diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-8n4g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-8n4g-async-1off-mxfp8-rollout.yaml index 204c96a5dc3..20396c4aa25 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-8n4g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-8n4g-async-1off-mxfp8-rollout.yaml @@ -10,6 +10,8 @@ policy: precision: "fp8" is_mx: true refit_prequantize: true + refit_batched_moe_shuffle: true + refit_cache_loader_routes: true quantization_ignored_layer_kws: - q_proj - k_proj diff --git a/nemo_rl/models/generation/vllm/config.py b/nemo_rl/models/generation/vllm/config.py index 858d0e007f5..dea33957110 100644 --- a/nemo_rl/models/generation/vllm/config.py +++ b/nemo_rl/models/generation/vllm/config.py @@ -41,6 +41,10 @@ class VllmSpecificArgs(TypedDict): # the vLLM worker then skips its per-refit re-quantization. Requires the # Megatron policy backend. refit_prequantize: NotRequired[bool] + # Batch MoE weight-layout transforms across experts during MXFP8 refit. + refit_batched_moe_shuffle: bool + # Cache and replay stable vLLM weight-loader routes across refits. + refit_cache_loader_routes: bool kv_cache_dtype: Literal["auto", "fp8", "fp8_e4m3"] enforce_eager: NotRequired[bool] enable_return_routed_experts: NotRequired[bool] @@ -175,6 +179,10 @@ def validate_vllm_quantization_config(config: VllmConfig) -> None: "policy.generation.vllm_cfg.refit_prequantize requires " "precision='fp8' and is_mx=true." ) + for field in ("refit_batched_moe_shuffle", "refit_cache_loader_routes"): + value = vllm_cfg.get(field) + if value is not None and not isinstance(value, bool): + raise ValueError(f"policy.generation.vllm_cfg.{field} must be a boolean.") def normalize_vllm_refit_config(config: VllmConfig) -> VllmRefitConfig | None: diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index dd6b3d7d43c..ffa277c47b7 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -62,6 +62,8 @@ class FP8Config: # Weights arrive from the trainer already MXFP8-quantized (E4M3 data plus # *_scale_from_checkpoint entries), so load_weights skips re-quantization. refit_prequantize: bool = False + refit_batched_moe_shuffle: bool = True + refit_cache_loader_routes: bool = False @dataclass() @@ -280,6 +282,8 @@ def init_fp8(vllm_cfg, model_name, model_parallel_size): "model_parallel_size": model_parallel_size, "kv_cache_dtype": kv_cache_dtype, "use_fp8_weights": use_fp8_weights, + "refit_batched_moe_shuffle": vllm_cfg["refit_batched_moe_shuffle"], + "refit_cache_loader_routes": vllm_cfg["refit_cache_loader_routes"], } if is_mx: fp8_config_kwargs["is_mx"] = True @@ -555,7 +559,11 @@ def load_weights(weights, model_runner): # module top would cycle through the nemo_rl generation package init. from nemo_rl.models.generation.vllm.vllm_backend import load_weights_maybe_cached - load_weights_maybe_cached(model, weights_quantized) + load_weights_maybe_cached( + model, + weights_quantized, + cache_loader_routes=global_fp8_config.refit_cache_loader_routes, + ) def cast_tensor_to_fp8_blockwise( @@ -1353,8 +1361,8 @@ def process_weights_after_loading_mxfp8_moe(self, layer: RoutedExperts) -> None: w13_weight = swap_w13_to_w31(w13_weight) w13_scale = swap_w13_to_w31(w13_scale) - # NRL_MXFP8_BATCHED_SHUFFLE=0 is the kill switch back to the per-expert path. - use_batched_shuffle = os.getenv("NRL_MXFP8_BATCHED_SHUFFLE", "1") != "0" + assert global_fp8_config is not None + use_batched_shuffle = global_fp8_config.refit_batched_moe_shuffle if use_batched_shuffle: ( w13_weight_shuffled, diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index eed389a936e..524c79fa5ad 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -218,18 +218,19 @@ def recorder(param, loaded_weight, *args, **kwargs): def load_weights_maybe_cached( - model: Any, weights: list[tuple[str, torch.Tensor]] + model: Any, + weights: list[tuple[str, torch.Tensor]], + *, + cache_loader_routes: bool, ) -> set[str]: - """model.load_weights with optional loader replay caching. + """Load weights, optionally replaying cached loader routes. - Opt-in via NRL_REFIT_CACHED_LOADERS=1 since model load_weights - implementations vary; the default is a plain model.load_weights call. Cached parameter identities are re-validated against named_parameters() on every call, so a process_weights_after_loading pass that replaces parameter objects drops the cache instead of loading into orphans. Returns the set of loaded weight names, mirroring model.load_weights. """ - if os.getenv("NRL_REFIT_CACHED_LOADERS") != "1": + if not cache_loader_routes: return model.load_weights(weights=weights) cache = getattr(model, "_nrl_refit_loader_cache", None) @@ -316,6 +317,7 @@ class VllmInternalWorkerExtension: _mtp_drafter_from_disk: bool = False _sparse_delta_applier: Any = None _nrl_named_parameters: dict[str, torch.nn.Parameter] + _refit_cache_loader_routes: bool = False def _get_named_parameters(self) -> dict[str, torch.nn.Parameter]: params = getattr(self, "_nrl_named_parameters", None) @@ -327,9 +329,11 @@ def _get_named_parameters(self) -> dict[str, torch.nn.Parameter]: def _load_full_hf_weights( self, policy_weights: list[tuple[str, torch.Tensor]] ) -> None: - # Refit optimization: replay cached weight-loader routing when - # NRL_REFIT_CACHED_LOADERS=1 (identity-validated), else plain load. - load_weights_maybe_cached(self.model_runner.model, policy_weights) + load_weights_maybe_cached( + self.model_runner.model, + policy_weights, + cache_loader_routes=self._refit_cache_loader_routes, + ) def _load_hf_weights(self, policy_weights: list[tuple[str, torch.Tensor]]) -> None: from nemo_rl.models.generation.vllm.quantization import fp8 @@ -470,6 +474,10 @@ def prepare_refit_info( from nemo_rl.models.generation.vllm.quantization import fp8 fp8.install_fp8_config(serialized_fp8_config) + self._refit_cache_loader_routes = bool( + fp8.global_fp8_config + and fp8.global_fp8_config.refit_cache_loader_routes + ) if not ( fp8.global_fp8_config is not None and fp8.global_fp8_config.is_mx diff --git a/tests/test_mxfp8_rollout_recipes.py b/tests/test_mxfp8_rollout_recipes.py index 65a9215df24..106bfdd65f0 100644 --- a/tests/test_mxfp8_rollout_recipes.py +++ b/tests/test_mxfp8_rollout_recipes.py @@ -129,6 +129,8 @@ def test_mxfp8_rollout_recipe_matrix(case_name: str, expected: dict) -> None: assert vllm_cfg["precision"] == "fp8" assert vllm_cfg["is_mx"] is True assert vllm_cfg["refit_prequantize"] is True + assert vllm_cfg["refit_batched_moe_shuffle"] is True + assert vllm_cfg["refit_cache_loader_routes"] is True assert config["policy"]["megatron_cfg"]["enabled"] is True expected_ignored = [ "q_proj", diff --git a/tests/unit/models/generation/test_vllm_config.py b/tests/unit/models/generation/test_vllm_config.py index bba7eceaec0..e6722fd6e24 100644 --- a/tests/unit/models/generation/test_vllm_config.py +++ b/tests/unit/models/generation/test_vllm_config.py @@ -90,6 +90,26 @@ def test_refit_prequantize_accepts_mxfp8() -> None: validate_vllm_quantization_config(generation_config) +@pytest.mark.parametrize( + "field", + ["refit_batched_moe_shuffle", "refit_cache_loader_routes"], +) +def test_refit_optimization_flags_must_be_boolean(field: str) -> None: + generation_config = cast( + VllmConfig, + { + "vllm_cfg": { + "precision": "fp8", + "is_mx": True, + field: "true", + } + }, + ) + + with pytest.raises(ValueError, match=rf"{field} must be a boolean"): + validate_vllm_quantization_config(generation_config) + + def test_refit_prequantize_validation_allows_omitted_vllm_cfg() -> None: generation_config = cast(VllmConfig, {"quant_cfg": None}) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 94a31564fbd..864a510a4d4 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -66,6 +66,8 @@ def test_init_fp8_uses_mxfp8_quantization_config(fp8_module, monkeypatch): "kv_cache_dtype": "auto", "async_engine": False, "is_mx": True, + "refit_batched_moe_shuffle": False, + "refit_cache_loader_routes": True, "use_deep_gemm": True, }, "dummy-model", @@ -85,6 +87,8 @@ def test_init_fp8_uses_mxfp8_quantization_config(fp8_module, monkeypatch): } assert applied_configs == [fp8.global_fp8_config] assert fp8.global_fp8_config.is_mx is True + assert fp8.global_fp8_config.refit_batched_moe_shuffle is False + assert fp8.global_fp8_config.refit_cache_loader_routes is True assert "VLLM_USE_DEEP_GEMM" not in fp8.os.environ assert "VLLM_USE_DEEP_GEMM_E8M0" not in fp8.os.environ @@ -167,6 +171,8 @@ def test_init_fp8_rejects_non_pow2_mxfp8_scales(fp8_module, monkeypatch, field, "kv_cache_dtype": "auto", "async_engine": False, "is_mx": True, + "refit_batched_moe_shuffle": True, + "refit_cache_loader_routes": False, field: False, }, "dummy-model", @@ -423,7 +429,7 @@ def fake_batched_shuffle( monkeypatch.setattr(fp8, "_shuffle_mxfp8_moe_batched", fake_batched_shuffle) monkeypatch.setattr(fp8_oracle, "make_fp8_moe_quant_config", fake_make_quant_config) monkeypatch.setattr(fp8_oracle, "make_fp8_moe_kernel", fake_make_kernel) - monkeypatch.delenv("NRL_MXFP8_BATCHED_SHUFFLE", raising=False) + fp8.global_fp8_config = fp8.FP8Config(refit_batched_moe_shuffle=True) layer = torch.nn.Module() layer.w13_weight = torch.nn.Parameter( diff --git a/tests/unit/models/generation/test_vllm_refit_loader.py b/tests/unit/models/generation/test_vllm_refit_loader.py index b2b9c660d1c..7878eac51a8 100644 --- a/tests/unit/models/generation/test_vllm_refit_loader.py +++ b/tests/unit/models/generation/test_vllm_refit_loader.py @@ -200,7 +200,6 @@ def load_weights(self, *, weights): loaded.add(name) return loaded - monkeypatch.setenv("NRL_REFIT_CACHED_LOADERS", "1") model = Model() first_expert = torch.tensor([1.0]) first_default = torch.tensor([2.0]) @@ -210,10 +209,12 @@ def load_weights(self, *, weights): assert load_weights_maybe_cached( model, [("expert", first_expert), ("default", first_default)], + cache_loader_routes=True, ) == {"expert", "default"} assert load_weights_maybe_cached( model, [("expert", second_expert), ("default", second_default)], + cache_loader_routes=True, ) == {"expert", "default"} cache = model._nrl_refit_loader_cache @@ -274,18 +275,21 @@ def load_weights(self, *, weights): self.local.weight_loader(self.local, weight) return {name for name, _weight in weights} - monkeypatch.setenv("NRL_REFIT_CACHED_LOADERS", "1") model = Model() first = torch.tensor([1.0]) second = torch.tensor([2.0]) - assert load_weights_maybe_cached(model, [("expert", first)]) == {"expert"} + assert load_weights_maybe_cached( + model, [("expert", first)], cache_loader_routes=True + ) == {"expert"} cache = model._nrl_refit_loader_cache old_local = model.local model.local = torch.nn.Parameter(torch.zeros(1), requires_grad=False) model.local.weight_loader = make_loader("replacement") - assert load_weights_maybe_cached(model, [("expert", second)]) == {"expert"} + assert load_weights_maybe_cached( + model, [("expert", second)], cache_loader_routes=True + ) == {"expert"} assert model.load_calls == [["expert"], ["expert"]] assert events == ["remote", "local", "remote", "replacement"] diff --git a/tests/unit/reference_configs/distillation_math.yaml b/tests/unit/reference_configs/distillation_math.yaml index 48b676707ca..cde9a427e53 100644 --- a/tests/unit/reference_configs/distillation_math.yaml +++ b/tests/unit/reference_configs/distillation_math.yaml @@ -195,6 +195,8 @@ policy: &POLICY_BASE vllm_cfg: async_engine: false precision: ${...precision} + refit_batched_moe_shuffle: true # Batch MoE layout transforms across experts during refit. + refit_cache_loader_routes: false # Replay stable vLLM loader routes across refits. kv_cache_dtype: "auto" tensor_parallel_size: 1 pipeline_parallel_size: 1 diff --git a/tests/unit/reference_configs/grpo_math_1B.yaml b/tests/unit/reference_configs/grpo_math_1B.yaml index 8539754cd64..157c6b91d1c 100644 --- a/tests/unit/reference_configs/grpo_math_1B.yaml +++ b/tests/unit/reference_configs/grpo_math_1B.yaml @@ -374,6 +374,8 @@ policy: precision: ${policy.precision} # MXFP8 + Megatron only: quantize on trainer and stream E4M3 values plus scales. refit_prequantize: false + refit_batched_moe_shuffle: true # Batch MoE layout transforms across experts during refit. + refit_cache_loader_routes: false # Replay stable vLLM loader routes across refits. kv_cache_dtype: "auto" tensor_parallel_size: 1 pipeline_parallel_size: 1 diff --git a/tests/unit/reference_configs/ppo_math_1B_megatron.yaml b/tests/unit/reference_configs/ppo_math_1B_megatron.yaml index 5b025b460b7..d539fb4b017 100644 --- a/tests/unit/reference_configs/ppo_math_1B_megatron.yaml +++ b/tests/unit/reference_configs/ppo_math_1B_megatron.yaml @@ -235,6 +235,8 @@ policy: vllm_cfg: async_engine: false precision: ${policy.precision} + refit_batched_moe_shuffle: true # Batch MoE layout transforms across experts during refit. + refit_cache_loader_routes: false # Replay stable vLLM loader routes across refits. kv_cache_dtype: "auto" tensor_parallel_size: 1 pipeline_parallel_size: 1 diff --git a/tests/unit/weight_sync/test_weight_synchronizer.py b/tests/unit/weight_sync/test_weight_synchronizer.py index 7380e9e3e31..508dab2bcff 100644 --- a/tests/unit/weight_sync/test_weight_synchronizer.py +++ b/tests/unit/weight_sync/test_weight_synchronizer.py @@ -471,6 +471,30 @@ def test_sync_weights_passes_kv_scales(self, mock_ray): call_kwargs = policy.broadcast_weights_for_collective.call_args assert call_kwargs.kwargs["kv_scales"] == kv_scales + @patch("nemo_rl.weight_sync.collective_weight_synchronizer.ray") + def test_sync_weights_forwards_fixed_buffer_size(self, mock_ray): + mock_ray.get.return_value = [True] + policy = _mock_policy() + gen = _mock_generation() + sync = CollectiveWeightSynchronizer( + policy, + gen, + _mock_cluster(), + _mock_cluster(), + refit_buffer_size_gb=1.5, + ) + + sync.sync_weights() + + expected_bytes = int(1.5 * 1024**3) + policy.broadcast_weights_for_collective.assert_called_once_with( + kv_scales=None, + buffer_size_bytes=expected_bytes, + ) + gen.update_weights_from_collective.assert_called_once_with( + buffer_size_bytes=expected_bytes + ) + @patch("nemo_rl.weight_sync.collective_weight_synchronizer.ray") def test_sync_weights_raises_on_failure(self, mock_ray): mock_ray.get.side_effect = [ From c4d453e1afa41bbcf57314aeb2167d4b38caf740 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sat, 1 Aug 2026 12:06:14 -0700 Subject: [PATCH 28/76] fix(refit): propagate runtime configuration Signed-off-by: seonjinn --- docs/guides/refit.md | 3 ++ examples/configs/evals/eval.yaml | 2 + examples/configs/evals/mmau.yaml | 2 + ...-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml | 2 + nemo_rl/algorithms/grpo.py | 22 ++++++++-- nemo_rl/models/generation/interfaces.py | 4 +- .../megatron/megatron_generation.py | 5 ++- .../generation/sglang/sglang_generation.py | 5 ++- .../generation/trtllm/trtllm_backend.py | 7 +++ .../generation/trtllm/trtllm_generation.py | 13 ++++-- .../generation/trtllm/trtllm_worker_async.py | 14 +++++- .../generation/vllm/quantization/fp8.py | 9 ++-- .../models/generation/vllm/vllm_backend.py | 23 ++++++---- .../models/generation/vllm/vllm_generation.py | 5 ++- nemo_rl/models/generation/vllm/vllm_worker.py | 18 ++++++-- .../generation/vllm/vllm_worker_async.py | 16 +++++-- .../models/generation/vllm/worker_utils.py | 28 ++++++++++++ nemo_rl/models/policy/interfaces.py | 4 +- nemo_rl/models/policy/lm_policy.py | 5 ++- .../policy/workers/dtensor_policy_worker.py | 5 ++- .../workers/dtensor_policy_worker_v2.py | 5 ++- .../policy/workers/megatron_policy_worker.py | 5 ++- nemo_rl/utils/packed_tensor.py | 32 ++++++++++++-- .../collective_weight_synchronizer.py | 27 ++++++++++-- nemo_rl/weight_sync/factory.py | 1 + tests/test_mxfp8_rollout_recipes.py | 15 +++++++ .../environments/test_code_environment.py | 2 + tests/unit/environments/test_retriever.py | 2 + tests/unit/experience/test_rollouts.py | 2 + .../models/generation/test_vllm_backend.py | 43 ++++++++++++++++++- .../generation/test_vllm_fp8_quantization.py | 8 ++-- .../models/generation/test_vllm_generation.py | 2 + .../generation/test_vllm_large_model.py | 2 + .../generation/test_vllm_quant_backend.py | 2 + .../generation/test_vllm_worker_helpers.py | 20 +++++++++ tests/unit/reference_configs/eval.yaml | 2 + tests/unit/utils/test_packed_tensor.py | 16 +++++++ .../weight_sync/test_weight_synchronizer.py | 2 + 38 files changed, 330 insertions(+), 50 deletions(-) diff --git a/docs/guides/refit.md b/docs/guides/refit.md index 3abcf974100..58bee51e015 100644 --- a/docs/guides/refit.md +++ b/docs/guides/refit.md @@ -40,6 +40,9 @@ GRPO and distillation setup paths; PPO currently requires colocated generation. | Option | Scope | Effect | |---|---|---| | `policy.generation.vllm_cfg.refit_prequantize` | Megatron training with MXFP8 vLLM rollout | Quantizes eligible weights on the trainer and transfers E4M3 values plus E8M0 scales. Requires `precision: fp8` and `is_mx: true`; sparse delta and NCCL Reshard do not support it. | +| `policy.generation.vllm_cfg.refit_batched_moe_shuffle` | MXFP8 vLLM rollout | Batches MoE weight-layout transforms across experts. Enabled by default. | +| `policy.generation.vllm_cfg.refit_cache_loader_routes` | vLLM refit | Replays identity-validated weight-loader routes after the first refit. Disabled by default because loader behavior is model-dependent. | +| `policy.refit_buffer_size_gb` | Colocated IPC/HTTP or non-colocated NCCL broadcast | Sets the packing threshold explicitly. For NCCL broadcast, the same byte value is sent to producer and consumer so their collective chunk boundaries match. | | `policy.refit_persistent_ipc_buffers` | Colocated CUDA-IPC refit | Reuses the two trainer staging buffers across refits. A fixed `refit_buffer_size_gb` gives stable memory use. | | `policy.megatron_cfg.refit_slim_offload_after` | Colocated Megatron refit | Avoids repeating grad-buffer offload and a second allocator cleanup after weights are transferred. | | `policy.megatron_cfg.pinned_reference_swap` | Megatron reference-policy logprobs | Keeps the CPU reference copy in pinned memory for faster host-to-device swaps, at the cost of additional pinned host memory. | diff --git a/examples/configs/evals/eval.yaml b/examples/configs/evals/eval.yaml index c492ac34cd1..05859cde775 100644 --- a/examples/configs/evals/eval.yaml +++ b/examples/configs/evals/eval.yaml @@ -19,6 +19,8 @@ generation: vllm_cfg: async_engine: false precision: "bfloat16" + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false tensor_parallel_size: 1 pipeline_parallel_size: 1 expert_parallel_size: 1 diff --git a/examples/configs/evals/mmau.yaml b/examples/configs/evals/mmau.yaml index e12c3ea0aec..fafe2a79155 100644 --- a/examples/configs/evals/mmau.yaml +++ b/examples/configs/evals/mmau.yaml @@ -18,6 +18,8 @@ generation: vllm_cfg: async_engine: false precision: "bfloat16" + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false tensor_parallel_size: 1 pipeline_parallel_size: 1 expert_parallel_size: 1 diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml index 4160e9b478a..67b5286a605 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml @@ -7,6 +7,8 @@ policy: precision: "fp8" is_mx: true refit_prequantize: true + refit_batched_moe_shuffle: true + refit_cache_loader_routes: true quantization_ignored_layer_kws: - q_proj - k_proj diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index e01b1cd55da..6ff5185ac60 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -2262,10 +2262,15 @@ def refit_policy_generation( with timer_context: # update weights update_success = False + configured_buffer_size_bytes = ( + None + if _refit_buffer_size_gb is None + else int(_refit_buffer_size_gb * 1024**3) + ) if colocated_inference: # get model param keys, which is grouped by size - if _refit_buffer_size_gb is not None: - buffer_size_bytes = int(_refit_buffer_size_gb * (1024**3)) + if configured_buffer_size_bytes is not None: + buffer_size_bytes = configured_buffer_size_bytes else: # Empirically sets ratio as 30% to maximize efficiency. # The remaining 70% is a necessary buffer reserved for the parameter all-gathering across the expert-parallelism dimension. @@ -2306,11 +2311,20 @@ def refit_policy_generation( ) if isinstance(policy_generation, MegatronGeneration): futures_train = policy.swap_weights_via_reshard(is_source=True) - else: + futures_inference = policy_generation.update_weights_from_collective() + elif configured_buffer_size_bytes is None: futures_train = policy.broadcast_weights_for_collective( kv_scales=kv_scales ) - futures_inference = policy_generation.update_weights_from_collective() + futures_inference = policy_generation.update_weights_from_collective() + else: + futures_train = policy.broadcast_weights_for_collective( + kv_scales=kv_scales, + buffer_size_bytes=configured_buffer_size_bytes, + ) + futures_inference = policy_generation.update_weights_from_collective( + buffer_size_bytes=configured_buffer_size_bytes + ) # wait for all futures to complete ray.get(futures_train) results = ray.get(futures_inference) diff --git a/nemo_rl/models/generation/interfaces.py b/nemo_rl/models/generation/interfaces.py index c82eeeac011..3eee36a921f 100644 --- a/nemo_rl/models/generation/interfaces.py +++ b/nemo_rl/models/generation/interfaces.py @@ -346,7 +346,9 @@ def update_weights_via_ipc_zmq(self) -> list[ray.ObjectRef]: """Update the model weights from the given IPC handles.""" raise NotImplementedError - def update_weights_from_collective(self) -> list[ray.ObjectRef]: + def update_weights_from_collective( + self, buffer_size_bytes: Optional[int] = None + ) -> list[ray.ObjectRef]: """Update the model weights from collective communication.""" raise NotImplementedError diff --git a/nemo_rl/models/generation/megatron/megatron_generation.py b/nemo_rl/models/generation/megatron/megatron_generation.py index c429afe0ab6..63b296b174c 100644 --- a/nemo_rl/models/generation/megatron/megatron_generation.py +++ b/nemo_rl/models/generation/megatron/megatron_generation.py @@ -158,8 +158,11 @@ def init_collective( refit_backend=refit_backend, ) - def update_weights_from_collective(self) -> list[ray.ObjectRef]: + def update_weights_from_collective( + self, buffer_size_bytes: Optional[int] = None + ) -> list[ray.ObjectRef]: """Receive updated weights from the training cluster via collective communication.""" + del buffer_size_bytes return self._policy.swap_weights_via_reshard(is_source=False) def generate( diff --git a/nemo_rl/models/generation/sglang/sglang_generation.py b/nemo_rl/models/generation/sglang/sglang_generation.py index 0aff8acee37..dacdab7e6a8 100644 --- a/nemo_rl/models/generation/sglang/sglang_generation.py +++ b/nemo_rl/models/generation/sglang/sglang_generation.py @@ -739,7 +739,10 @@ def prepare_refit_info( def update_weights_via_ipc_zmq(self) -> list[ray.ObjectRef]: return [] - def update_weights_from_collective(self) -> list[ray.ObjectRef]: + def update_weights_from_collective( + self, buffer_size_bytes: int | None = None + ) -> list[ray.ObjectRef]: + del buffer_size_bytes return [] def prepare_for_generation(self, *args: Any, **kwargs: Any) -> bool: diff --git a/nemo_rl/models/generation/trtllm/trtllm_backend.py b/nemo_rl/models/generation/trtllm/trtllm_backend.py index 184b8cb4f21..f5b7737d131 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_backend.py +++ b/nemo_rl/models/generation/trtllm/trtllm_backend.py @@ -121,6 +121,7 @@ def update_weights_from_collective( *, drain: bool = True, recompute_kv: bool = False, + buffer_size_bytes: int | None = None, ) -> bool: """Receive weights via NCCL broadcast and update model parameters. @@ -163,11 +164,17 @@ def load_model_weight_func(weight_list): module, "_weights_removed", False ): module.pre_reload_weights() + consumer_kwargs = ( + {} + if buffer_size_bytes is None + else {"buffer_size_bytes": buffer_size_bytes} + ) packed_broadcast_consumer( iterator=iter(self.state_dict_info.items()), group=self.model_update_group, src=0, post_unpack_func=load_model_weight_func, + **consumer_kwargs, ) _call_model_loader_hook_if_available( model_engine.model_loader, "finalize_update_weights" diff --git a/nemo_rl/models/generation/trtllm/trtllm_generation.py b/nemo_rl/models/generation/trtllm/trtllm_generation.py index 89f153b48a8..9884c0233a4 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_generation.py +++ b/nemo_rl/models/generation/trtllm/trtllm_generation.py @@ -468,17 +468,24 @@ def stop_gpu_profiling(self) -> None: ) ray.get(futures) - def update_weights_from_collective(self) -> list[ray.ObjectRef]: + def update_weights_from_collective( + self, buffer_size_bytes: int | None = None + ) -> list[ray.ObjectRef]: if not self.worker_group or not self.worker_group.workers: raise RuntimeError("Worker group not initialised") trtllm_cfg = self.cfg["trtllm_cfg"] in_flight = bool(trtllm_cfg.get("in_flight_weight_updates")) recompute_kv = bool(trtllm_cfg.get("recompute_kv_cache_after_weight_updates")) + update_kwargs: dict[str, bool | int] = { + "drain": not in_flight, + "recompute_kv": recompute_kv, + } + if buffer_size_bytes is not None: + update_kwargs["buffer_size_bytes"] = buffer_size_bytes return self.worker_group.run_all_workers_single_data( "update_weights_from_collective_async", run_rank_0_only_axes=["tensor_parallel"], - drain=not in_flight, - recompute_kv=recompute_kv, + **update_kwargs, ) def update_weights_via_ipc_zmq(self) -> list[ray.ObjectRef]: diff --git a/nemo_rl/models/generation/trtllm/trtllm_worker_async.py b/nemo_rl/models/generation/trtllm/trtllm_worker_async.py index 5f4b1002c99..dff17038436 100644 --- a/nemo_rl/models/generation/trtllm/trtllm_worker_async.py +++ b/nemo_rl/models/generation/trtllm/trtllm_worker_async.py @@ -342,7 +342,11 @@ async def prepare_refit_info_async(self, state_dict_info: dict[str, Any]) -> Non await self.llm.collective_rpc("prepare_refit_info", args=(state_dict_info,)) async def update_weights_from_collective_async( - self, *, drain: bool = True, recompute_kv: bool = False + self, + *, + drain: bool = True, + recompute_kv: bool = False, + buffer_size_bytes: int | None = None, ) -> bool: """Async version of ``update_weights_from_collective``. @@ -357,9 +361,15 @@ async def update_weights_from_collective_async( """ assert self.llm is not None try: + update_kwargs: dict[str, bool | int] = { + "drain": drain, + "recompute_kv": recompute_kv, + } + if buffer_size_bytes is not None: + update_kwargs["buffer_size_bytes"] = buffer_size_bytes results = await self.llm.collective_rpc( "update_weights_from_collective", - kwargs={"drain": drain, "recompute_kv": recompute_kv}, + kwargs=update_kwargs, ) worker_result = results[0] if results else True if not worker_result: diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index ffa277c47b7..937b2c4cf29 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -33,6 +33,9 @@ from nemo_rl.models.generation.vllm.quantization.mxfp8_utils import ( pad_flashinfer_scale_k, ) +from nemo_rl.models.generation.vllm.worker_utils import ( + refit_cache_loader_routes_enabled, +) logger = init_logger(__name__) @@ -63,7 +66,6 @@ class FP8Config: # *_scale_from_checkpoint entries), so load_weights skips re-quantization. refit_prequantize: bool = False refit_batched_moe_shuffle: bool = True - refit_cache_loader_routes: bool = False @dataclass() @@ -283,7 +285,6 @@ def init_fp8(vllm_cfg, model_name, model_parallel_size): "kv_cache_dtype": kv_cache_dtype, "use_fp8_weights": use_fp8_weights, "refit_batched_moe_shuffle": vllm_cfg["refit_batched_moe_shuffle"], - "refit_cache_loader_routes": vllm_cfg["refit_cache_loader_routes"], } if is_mx: fp8_config_kwargs["is_mx"] = True @@ -562,7 +563,9 @@ def load_weights(weights, model_runner): load_weights_maybe_cached( model, weights_quantized, - cache_loader_routes=global_fp8_config.refit_cache_loader_routes, + cache_loader_routes=refit_cache_loader_routes_enabled( + model_runner.vllm_config + ), ) diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index 524c79fa5ad..e6671fc8b24 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -33,6 +33,9 @@ calculate_aligned_size, rebuild_cuda_tensor_from_ipc, ) +from nemo_rl.models.generation.vllm.worker_utils import ( + refit_cache_loader_routes_enabled, +) from nemo_rl.utils.nsys import wrap_with_nvtx_name from nemo_rl.utils.packed_tensor import packed_broadcast_consumer from nemo_rl.weight_sync.nccl_reshard_utils import ( @@ -317,8 +320,6 @@ class VllmInternalWorkerExtension: _mtp_drafter_from_disk: bool = False _sparse_delta_applier: Any = None _nrl_named_parameters: dict[str, torch.nn.Parameter] - _refit_cache_loader_routes: bool = False - def _get_named_parameters(self) -> dict[str, torch.nn.Parameter]: params = getattr(self, "_nrl_named_parameters", None) if params is None: @@ -332,7 +333,9 @@ def _load_full_hf_weights( load_weights_maybe_cached( self.model_runner.model, policy_weights, - cache_loader_routes=self._refit_cache_loader_routes, + cache_loader_routes=refit_cache_loader_routes_enabled( + self.model_runner.vllm_config + ), ) def _load_hf_weights(self, policy_weights: list[tuple[str, torch.Tensor]]) -> None: @@ -474,10 +477,6 @@ def prepare_refit_info( from nemo_rl.models.generation.vllm.quantization import fp8 fp8.install_fp8_config(serialized_fp8_config) - self._refit_cache_loader_routes = bool( - fp8.global_fp8_config - and fp8.global_fp8_config.refit_cache_loader_routes - ) if not ( fp8.global_fp8_config is not None and fp8.global_fp8_config.is_mx @@ -893,7 +892,9 @@ def update_weights_via_ipc_zmq(self) -> bool: @wrap_with_nvtx_name( "vllm_internal_worker_extension/update_weights_from_collective" ) - def update_weights_from_collective(self) -> bool: + def update_weights_from_collective( + self, buffer_size_bytes: Optional[int] = None + ) -> bool: """Update the model weights from collective communication.""" assert self.state_dict_info is not None, ( "state_dict_info is not prepared. " @@ -902,11 +903,17 @@ def update_weights_from_collective(self) -> bool: try: with self._weight_update_lifecycle("collective") as finalize: + consumer_kwargs = ( + {} + if buffer_size_bytes is None + else {"buffer_size_bytes": buffer_size_bytes} + ) packed_broadcast_consumer( iterator=iter(self.state_dict_info.items()), group=self.model_update_group, src=0, post_unpack_func=self._load_weights, + **consumer_kwargs, ) finalize() diff --git a/nemo_rl/models/generation/vllm/vllm_generation.py b/nemo_rl/models/generation/vllm/vllm_generation.py index 0343ae970ac..a14491c002b 100644 --- a/nemo_rl/models/generation/vllm/vllm_generation.py +++ b/nemo_rl/models/generation/vllm/vllm_generation.py @@ -977,7 +977,9 @@ def update_weights_via_ipc_zmq(self) -> list[ray.ObjectRef]: # this function should co-work with lm_policy, so we should wait for all futures to complete outside return futures - def update_weights_from_collective(self) -> list[ray.ObjectRef]: + def update_weights_from_collective( + self, buffer_size_bytes: Optional[int] = None + ) -> list[ray.ObjectRef]: """Update weights of the policy using collective communication.""" if not self.worker_group or not self.worker_group.workers: raise RuntimeError("Worker group is not initialized") @@ -993,6 +995,7 @@ def update_weights_from_collective(self) -> list[ray.ObjectRef]: futures = self.worker_group.run_all_workers_single_data( method_name, run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], + buffer_size_bytes=buffer_size_bytes, ) # this function should co-work with lm_policy, so we should wait for all futures to complete outside diff --git a/nemo_rl/models/generation/vllm/vllm_worker.py b/nemo_rl/models/generation/vllm/vllm_worker.py index 5186bea6811..9866a88bedb 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker.py +++ b/nemo_rl/models/generation/vllm/vllm_worker.py @@ -50,6 +50,7 @@ pad_and_align_routed_expert_indices, ) from nemo_rl.models.generation.vllm.worker_utils import ( + configure_refit_runtime, resolve_data_parallel_local_rank, resolve_distributed_executor_backend, ) @@ -355,6 +356,7 @@ def _load_model(self, bundle_indices, seed): "please run at least once with the environment variable NRL_FORCE_REBUILD_VENVS=true set to force the rebuild of the environment." ) vllm_kwargs: dict[str, Any] = copy.deepcopy(self.cfg.get("vllm_kwargs", {})) + configure_refit_runtime(self.cfg["vllm_cfg"], vllm_kwargs) checkpoint_engine_config = checkpoint_engine_refit_config(self.cfg) if checkpoint_engine_config is not None: from nemo_rl.models.generation.vllm.checkpoint_engine import ( @@ -1140,7 +1142,9 @@ def update_weights_via_ipc_zmq(self) -> bool: return False @wrap_with_nvtx_name("vllm_genertion_worker/update_weights_from_collective") - def update_weights_from_collective(self) -> bool: + def update_weights_from_collective( + self, buffer_size_bytes: Optional[int] = None + ) -> bool: """Update the model weights from collective communication.""" try: assert self.llm is not None, ( @@ -1152,9 +1156,15 @@ def update_weights_from_collective(self) -> bool: "update_weights_from_collective can only be used with async_engine=False. Use update_weights_from_collective_async instead." ) - result_or_coro = self.llm.collective_rpc( - "update_weights_from_collective", args=tuple() - ) + if buffer_size_bytes is None: + result_or_coro = self.llm.collective_rpc( + "update_weights_from_collective", args=tuple() + ) + else: + result_or_coro = self.llm.collective_rpc( + "update_weights_from_collective", + kwargs={"buffer_size_bytes": buffer_size_bytes}, + ) worker_results = cast(list[bool], result_or_coro) if not worker_results or not all(worker_results): diff --git a/nemo_rl/models/generation/vllm/vllm_worker_async.py b/nemo_rl/models/generation/vllm/vllm_worker_async.py index 86c3ad7eac0..520dec843bd 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker_async.py +++ b/nemo_rl/models/generation/vllm/vllm_worker_async.py @@ -1354,7 +1354,9 @@ async def update_weights_via_ipc_zmq_async( traceback.print_exc() return False - async def update_weights_from_collective_async(self) -> bool: + async def update_weights_from_collective_async( + self, buffer_size_bytes: Optional[int] = None + ) -> bool: """Async version of update_weights_from_collective.""" try: assert self.llm is not None, ( @@ -1366,9 +1368,15 @@ async def update_weights_from_collective_async(self) -> bool: "update_weights_from_collective_async can only be used with async_engine=True. Use update_weights_from_collective instead." ) - result_or_coro = await self.llm.collective_rpc( - "update_weights_from_collective", args=tuple() - ) + if buffer_size_bytes is None: + result_or_coro = await self.llm.collective_rpc( + "update_weights_from_collective", args=tuple() + ) + else: + result_or_coro = await self.llm.collective_rpc( + "update_weights_from_collective", + kwargs={"buffer_size_bytes": buffer_size_bytes}, + ) if asyncio.iscoroutine(result_or_coro): worker_results = await result_or_coro diff --git a/nemo_rl/models/generation/vllm/worker_utils.py b/nemo_rl/models/generation/vllm/worker_utils.py index 831fbde0dd5..70cf996e0d1 100644 --- a/nemo_rl/models/generation/vllm/worker_utils.py +++ b/nemo_rl/models/generation/vllm/worker_utils.py @@ -12,6 +12,34 @@ # See the License for the specific language governing permissions and # limitations under the License. +from collections.abc import Mapping +from typing import Any + + +_REFIT_CACHE_LOADER_ROUTES_KEY = "nemo_rl_refit_cache_loader_routes" + + +def configure_refit_runtime( + vllm_cfg: Mapping[str, Any], vllm_kwargs: dict[str, Any] +) -> None: + """Forward NeMo-RL refit options through vLLM's worker config.""" + additional_config = dict(vllm_kwargs.get("additional_config") or {}) + additional_config[_REFIT_CACHE_LOADER_ROUTES_KEY] = vllm_cfg[ + "refit_cache_loader_routes" + ] + vllm_kwargs["additional_config"] = additional_config + + +def refit_cache_loader_routes_enabled(vllm_config: Any) -> bool: + """Return the configured loader-route cache setting in a vLLM worker.""" + additional_config = getattr(vllm_config, "additional_config", None) or {} + if _REFIT_CACHE_LOADER_ROUTES_KEY not in additional_config: + return False + value = additional_config[_REFIT_CACHE_LOADER_ROUTES_KEY] + if not isinstance(value, bool): + raise TypeError(f"{_REFIT_CACHE_LOADER_ROUTES_KEY} must be a boolean") + return value + def resolve_distributed_executor_backend( tensor_parallel_size: int, diff --git a/nemo_rl/models/policy/interfaces.py b/nemo_rl/models/policy/interfaces.py index 6c1c1b5aeaf..abdf37955b7 100644 --- a/nemo_rl/models/policy/interfaces.py +++ b/nemo_rl/models/policy/interfaces.py @@ -237,7 +237,9 @@ def set_rollout_num_gpus_per_engine(self, num_gpus_per_engine: int) -> None: @abstractmethod def broadcast_weights_for_collective( - self, kv_scales: Optional[dict[str, float]] = None + self, + kv_scales: Optional[dict[str, float]] = None, + buffer_size_bytes: Optional[int] = None, ) -> list[ray.ObjectRef]: pass diff --git a/nemo_rl/models/policy/lm_policy.py b/nemo_rl/models/policy/lm_policy.py index 20ca7abe715..bf04950c311 100644 --- a/nemo_rl/models/policy/lm_policy.py +++ b/nemo_rl/models/policy/lm_policy.py @@ -1080,12 +1080,15 @@ def set_rollout_num_gpus_per_engine(self, num_gpus_per_engine: int) -> None: ) def broadcast_weights_for_collective( - self, kv_scales: Optional[dict[str, float]] = None + self, + kv_scales: Optional[dict[str, float]] = None, + buffer_size_bytes: Optional[int] = None, ) -> list[ray.ObjectRef]: """Broadcast the weights for collective communication.""" futures = self.worker_group.run_all_workers_single_data( "broadcast_weights_for_collective", kv_scales=kv_scales, + buffer_size_bytes=buffer_size_bytes, ) # this function should co-work with vllm, so we should wait for all futures to complete outside return futures diff --git a/nemo_rl/models/policy/workers/dtensor_policy_worker.py b/nemo_rl/models/policy/workers/dtensor_policy_worker.py index 1803c371852..94468c71066 100644 --- a/nemo_rl/models/policy/workers/dtensor_policy_worker.py +++ b/nemo_rl/models/policy/workers/dtensor_policy_worker.py @@ -1902,7 +1902,9 @@ def _checkpoint_engine_params( @torch.no_grad() def broadcast_weights_for_collective( - self, kv_scales: Optional[dict[str, float]] = None + self, + kv_scales: Optional[dict[str, float]] = None, + buffer_size_bytes: Optional[int] = None, ) -> None: """Broadcast the weights for collective communication.""" if kv_scales is not None: @@ -1932,6 +1934,7 @@ def _dtensor_post_iter_func(tensor, dtype): group=self.model_update_group, src=0, post_iter_func=dtensor_post_iter_func, + buffer_size_bytes=buffer_size_bytes, ) # Manually move model to cpu for cpu offload case diff --git a/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py b/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py index 93856922579..30bacf933d5 100644 --- a/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py +++ b/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py @@ -1189,7 +1189,9 @@ def _checkpoint_engine_params( @torch.no_grad() def broadcast_weights_for_collective( - self, kv_scales: Optional[dict[str, float]] = None + self, + kv_scales: Optional[dict[str, float]] = None, + buffer_size_bytes: Optional[int] = None, ) -> None: """Broadcast the weights for collective communication.""" if kv_scales is not None: @@ -1213,6 +1215,7 @@ def broadcast_weights_for_collective( group=self.model_update_group, src=0, post_iter_func=dtensor_post_iter_func, + buffer_size_bytes=buffer_size_bytes, ) # Manually move model to cpu for cpu offload case diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 9ef392ad5b9..6460aafaa4e 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -2266,7 +2266,9 @@ def stream_weights_via_ipc_zmq( @torch.no_grad() def broadcast_weights_for_collective( - self, kv_scales: Optional[dict[str, float]] = None + self, + kv_scales: Optional[dict[str, float]] = None, + buffer_size_bytes: Optional[int] = None, ) -> None: """Broadcast the weights for collective communication.""" # param_iterator will return (name, tensor), we only need tensor. @@ -2275,6 +2277,7 @@ def broadcast_weights_for_collective( group=self.model_update_group, src=0, post_iter_func=lambda x: x[1], + buffer_size_bytes=buffer_size_bytes, ) def _build_layer_to_pp_stage( diff --git a/nemo_rl/utils/packed_tensor.py b/nemo_rl/utils/packed_tensor.py index 01f58c55a32..61e77c04bcb 100644 --- a/nemo_rl/utils/packed_tensor.py +++ b/nemo_rl/utils/packed_tensor.py @@ -36,7 +36,22 @@ def get_num_buffers(): return int(os.getenv("NRL_REFIT_NUM_BUFFERS", "2")) -def packed_broadcast_producer(iterator, group, src, post_iter_func): +def _resolve_target_packed_tensor_size(buffer_size_bytes: int | None) -> int: + if buffer_size_bytes is None: + return get_target_packed_tensor_size() + if buffer_size_bytes <= 0: + raise ValueError("buffer_size_bytes must be > 0") + return buffer_size_bytes + + +def packed_broadcast_producer( + iterator, + group, + src, + post_iter_func, + *, + buffer_size_bytes: int | None = None, +): """Broadcast a list of tensors in a packed manner. Args: @@ -44,12 +59,13 @@ def packed_broadcast_producer(iterator, group, src, post_iter_func): group: process group (vllm PyNcclCommunicator) src: source rank (0 in current implementation) post_iter_func: function to apply to each tensor before packing, should return a tensor + buffer_size_bytes: Optional explicit packing threshold. Returns: None """ - target_packed_tensor_size = get_target_packed_tensor_size() + target_packed_tensor_size = _resolve_target_packed_tensor_size(buffer_size_bytes) num_buffers = get_num_buffers() streams = [torch.cuda.Stream() for _ in range(num_buffers)] @@ -109,7 +125,14 @@ def packed_broadcast_producer(iterator, group, src, post_iter_func): s.synchronize() -def packed_broadcast_consumer(iterator, group, src, post_unpack_func): +def packed_broadcast_consumer( + iterator, + group, + src, + post_unpack_func, + *, + buffer_size_bytes: int | None = None, +): """Consume a packed tensor and unpack it into a list of tensors. Args: @@ -117,6 +140,7 @@ def packed_broadcast_consumer(iterator, group, src, post_unpack_func): group: process group (vllm PyNcclCommunicator) src: source rank (0 in current implementation) post_unpack_func: function to apply to each tensor after unpacking + buffer_size_bytes: Optional explicit packing threshold. Returns: None @@ -151,7 +175,7 @@ def unpack_tensor( return unpacked_list - target_packed_tensor_size = get_target_packed_tensor_size() + target_packed_tensor_size = _resolve_target_packed_tensor_size(buffer_size_bytes) num_buffers = get_num_buffers() streams = [torch.cuda.Stream() for _ in range(num_buffers)] diff --git a/nemo_rl/weight_sync/collective_weight_synchronizer.py b/nemo_rl/weight_sync/collective_weight_synchronizer.py index 13f23a1ee17..afb56527cec 100644 --- a/nemo_rl/weight_sync/collective_weight_synchronizer.py +++ b/nemo_rl/weight_sync/collective_weight_synchronizer.py @@ -52,6 +52,8 @@ class CollectiveWeightSynchronizer(WeightSynchronizer): train_cluster: RayVirtualCluster for the training workers, used to obtain the master address/port and world size for collective init. inference_cluster: RayVirtualCluster for the inference workers. + refit_buffer_size_gb: Optional fixed packing threshold shared by the + collective producer and consumer. """ def __init__( @@ -60,11 +62,19 @@ def __init__( generation: Any, train_cluster: Any, inference_cluster: Any, + refit_buffer_size_gb: float | int | None = None, ): self._policy = policy self._generation = generation self._train_cluster = train_cluster self._inference_cluster = inference_cluster + if refit_buffer_size_gb is not None and refit_buffer_size_gb <= 0: + raise ValueError("refit_buffer_size_gb must be > 0") + self._buffer_size_bytes = ( + None + if refit_buffer_size_gb is None + else int(refit_buffer_size_gb * 1024**3) + ) self._stale = True def sync_weights( @@ -79,10 +89,19 @@ def sync_weights( else nullcontext() ) with timer_context: - futures_train = self._policy.broadcast_weights_for_collective( - kv_scales=kv_scales - ) - futures_inference = self._generation.update_weights_from_collective() + if self._buffer_size_bytes is None: + futures_train = self._policy.broadcast_weights_for_collective( + kv_scales=kv_scales + ) + futures_inference = self._generation.update_weights_from_collective() + else: + futures_train = self._policy.broadcast_weights_for_collective( + kv_scales=kv_scales, + buffer_size_bytes=self._buffer_size_bytes, + ) + futures_inference = self._generation.update_weights_from_collective( + buffer_size_bytes=self._buffer_size_bytes + ) ray.get(futures_train) results = ray.get(futures_inference) diff --git a/nemo_rl/weight_sync/factory.py b/nemo_rl/weight_sync/factory.py index b76abf9933a..5a223491b0d 100644 --- a/nemo_rl/weight_sync/factory.py +++ b/nemo_rl/weight_sync/factory.py @@ -125,6 +125,7 @@ def create_weight_synchronizer( generation=generation, train_cluster=train_cluster, inference_cluster=inference_cluster, + refit_buffer_size_gb=refit_buffer_size_gb, ) if generation_backend == SGLANG_BACKEND: diff --git a/tests/test_mxfp8_rollout_recipes.py b/tests/test_mxfp8_rollout_recipes.py index 106bfdd65f0..dc8161a6eff 100644 --- a/tests/test_mxfp8_rollout_recipes.py +++ b/tests/test_mxfp8_rollout_recipes.py @@ -155,6 +155,21 @@ def test_mxfp8_rollout_recipe_matrix(case_name: str, expected: dict) -> None: ) +@pytest.mark.parametrize( + "config_path", + sorted(PERF_CONFIG_DIR.glob("*mxfp8-rollout.yaml")), + ids=lambda path: path.stem, +) +def test_all_mxfp8_rollout_recipes_enable_refit_optimizations( + config_path: Path, +) -> None: + config = _load_resolved_yaml(config_path) + vllm_cfg = config["policy"]["generation"]["vllm_cfg"] + + assert vllm_cfg["refit_batched_moe_shuffle"] is True + assert vllm_cfg["refit_cache_loader_routes"] is True + + def test_mxfp8_rollout_recipes_are_in_gb200_performance_suite() -> None: suite_text = GB200_SUITE.read_text(encoding="utf-8") diff --git a/tests/unit/environments/test_code_environment.py b/tests/unit/environments/test_code_environment.py index d32550aba1e..766e8376cb8 100644 --- a/tests/unit/environments/test_code_environment.py +++ b/tests/unit/environments/test_code_environment.py @@ -51,6 +51,8 @@ "vllm_cfg": { "async_engine": False, "precision": "bfloat16", + "refit_batched_moe_shuffle": True, + "refit_cache_loader_routes": False, "tensor_parallel_size": 1, "pipeline_parallel_size": 1, "expert_parallel_size": 1, diff --git a/tests/unit/environments/test_retriever.py b/tests/unit/environments/test_retriever.py index c9413e67590..53b6a995c35 100644 --- a/tests/unit/environments/test_retriever.py +++ b/tests/unit/environments/test_retriever.py @@ -50,6 +50,8 @@ "vllm_cfg": { "async_engine": False, "precision": "bfloat16", + "refit_batched_moe_shuffle": True, + "refit_cache_loader_routes": False, "tensor_parallel_size": 1, "pipeline_parallel_size": 1, "expert_parallel_size": 1, diff --git a/tests/unit/experience/test_rollouts.py b/tests/unit/experience/test_rollouts.py index 038384a5958..21aa17d0961 100644 --- a/tests/unit/experience/test_rollouts.py +++ b/tests/unit/experience/test_rollouts.py @@ -362,6 +362,8 @@ def initial_multi_step_calculator_batch(rollout_tokenizer): "vllm_cfg": { "async_engine": False, "precision": "bfloat16", + "refit_batched_moe_shuffle": True, + "refit_cache_loader_routes": False, "tensor_parallel_size": 1, "pipeline_parallel_size": 1, "expert_parallel_size": 1, diff --git a/tests/unit/models/generation/test_vllm_backend.py b/tests/unit/models/generation/test_vllm_backend.py index 05c1a5eecd2..ca62b68f0ba 100644 --- a/tests/unit/models/generation/test_vllm_backend.py +++ b/tests/unit/models/generation/test_vllm_backend.py @@ -272,8 +272,11 @@ def load_weights(weights): call_order.append("load") assert weights == [("model.weight", "weight-value")] - def packed_broadcast_consumer(iterator, group, src, post_unpack_func): + def packed_broadcast_consumer( + iterator, group, src, post_unpack_func, *, buffer_size_bytes + ): call_order.append("broadcast") + assert buffer_size_bytes == 1024 assert list(iterator) == [("model.weight", expected_state_info)] assert group is ext.model_update_group assert src == 0 @@ -290,7 +293,7 @@ def packed_broadcast_consumer(iterator, group, src, post_unpack_func): lambda: call_order.append("empty_cache"), ) - assert ext.update_weights_from_collective() is True + assert ext.update_weights_from_collective(buffer_size_bytes=1024) is True expected_process_calls = [(ext.model_runner.model, ext.model_config, ext.device)] expected_call_order = [ @@ -365,6 +368,21 @@ def test_sync_weight_updates_check_every_internal_worker( assert getattr(worker, method_name)() is expected +@pytest.mark.vllm +def test_sync_collective_update_forwards_buffer_size(): + from nemo_rl.models.generation.vllm.vllm_worker import VllmGenerationWorkerImpl + + worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) + worker.cfg = {"vllm_cfg": {"async_engine": False}} + worker.llm = SimpleNamespace(collective_rpc=MagicMock(return_value=[True])) + + assert worker.update_weights_from_collective(buffer_size_bytes=1024) is True + worker.llm.collective_rpc.assert_called_once_with( + "update_weights_from_collective", + kwargs={"buffer_size_bytes": 1024}, + ) + + @pytest.mark.vllm @pytest.mark.asyncio @pytest.mark.parametrize( @@ -389,6 +407,27 @@ async def test_async_weight_updates_check_every_internal_worker( assert await getattr(worker, method_name)() is expected +@pytest.mark.vllm +@pytest.mark.asyncio +async def test_async_collective_update_forwards_buffer_size(): + from nemo_rl.models.generation.vllm.vllm_worker_async import ( + VllmAsyncGenerationWorkerImpl, + ) + + worker = VllmAsyncGenerationWorkerImpl.__new__(VllmAsyncGenerationWorkerImpl) + worker.cfg = {"vllm_cfg": {"async_engine": True}} + worker.llm = SimpleNamespace(collective_rpc=AsyncMock(return_value=[True])) + + assert ( + await worker.update_weights_from_collective_async(buffer_size_bytes=1024) + is True + ) + worker.llm.collective_rpc.assert_awaited_once_with( + "update_weights_from_collective", + kwargs={"buffer_size_bytes": 1024}, + ) + + @pytest.mark.vllm def test_update_weights_via_ipc_acks_manifest_error_and_returns_false(monkeypatch): from nemo_rl.models.generation.vllm import vllm_backend diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 864a510a4d4..33c1f2e1345 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -88,7 +88,6 @@ def test_init_fp8_uses_mxfp8_quantization_config(fp8_module, monkeypatch): assert applied_configs == [fp8.global_fp8_config] assert fp8.global_fp8_config.is_mx is True assert fp8.global_fp8_config.refit_batched_moe_shuffle is False - assert fp8.global_fp8_config.refit_cache_loader_routes is True assert "VLLM_USE_DEEP_GEMM" not in fp8.os.environ assert "VLLM_USE_DEEP_GEMM_E8M0" not in fp8.os.environ @@ -298,7 +297,7 @@ def test_load_weights_preserves_prequantized_mxfp8_and_clamps_scales( monkeypatch.setattr( vllm_backend, "load_weights_maybe_cached", - lambda model, weights: loaded.extend(weights), + lambda model, weights, *, cache_loader_routes: loaded.extend(weights), ) model = object() @@ -308,7 +307,10 @@ def test_load_weights_preserves_prequantized_mxfp8_and_clamps_scales( ("model.prequantized.weight", prequantized), ("model.receiver.weight", receiver_quantized), ], - types.SimpleNamespace(model=model), + types.SimpleNamespace( + model=model, + vllm_config=types.SimpleNamespace(additional_config={}), + ), ) assert loaded[0][0] == "model.native" diff --git a/tests/unit/models/generation/test_vllm_generation.py b/tests/unit/models/generation/test_vllm_generation.py index d49c54b00e5..1f0a0286da7 100644 --- a/tests/unit/models/generation/test_vllm_generation.py +++ b/tests/unit/models/generation/test_vllm_generation.py @@ -69,6 +69,8 @@ "stop_strings": None, "vllm_cfg": { "precision": "bfloat16", + "refit_batched_moe_shuffle": True, + "refit_cache_loader_routes": False, "tensor_parallel_size": 1, "pipeline_parallel_size": 1, "expert_parallel_size": 1, diff --git a/tests/unit/models/generation/test_vllm_large_model.py b/tests/unit/models/generation/test_vllm_large_model.py index 89eaece234c..d500b5828dd 100644 --- a/tests/unit/models/generation/test_vllm_large_model.py +++ b/tests/unit/models/generation/test_vllm_large_model.py @@ -42,6 +42,8 @@ "stop_strings": None, "vllm_cfg": { "precision": "bfloat16", + "refit_batched_moe_shuffle": True, + "refit_cache_loader_routes": False, "tensor_parallel_size": 8, "pipeline_parallel_size": 2, "expert_parallel_size": 1, diff --git a/tests/unit/models/generation/test_vllm_quant_backend.py b/tests/unit/models/generation/test_vllm_quant_backend.py index 6b8bd6487c4..fa0cce8248f 100644 --- a/tests/unit/models/generation/test_vllm_quant_backend.py +++ b/tests/unit/models/generation/test_vllm_quant_backend.py @@ -62,6 +62,8 @@ def _make_vllm_config(tokenizer, *, async_engine=False, is_eval=True): "quant_cfg": _QUANT_CFG, "vllm_cfg": { "precision": "bfloat16", + "refit_batched_moe_shuffle": True, + "refit_cache_loader_routes": False, "tensor_parallel_size": 1, "pipeline_parallel_size": 1, "expert_parallel_size": 1, diff --git a/tests/unit/models/generation/test_vllm_worker_helpers.py b/tests/unit/models/generation/test_vllm_worker_helpers.py index 29d61934447..e600ff82f1a 100644 --- a/tests/unit/models/generation/test_vllm_worker_helpers.py +++ b/tests/unit/models/generation/test_vllm_worker_helpers.py @@ -14,14 +14,34 @@ """Tests for vLLM worker helper functions.""" +from types import SimpleNamespace + import pytest from nemo_rl.models.generation.vllm.worker_utils import ( + configure_refit_runtime, + refit_cache_loader_routes_enabled, resolve_data_parallel_local_rank, resolve_distributed_executor_backend, ) +@pytest.mark.parametrize("enabled", [False, True]) +def test_refit_loader_cache_round_trips_through_additional_config(enabled): + vllm_kwargs = {"additional_config": {"existing": "value"}} + + configure_refit_runtime( + {"refit_cache_loader_routes": enabled}, + vllm_kwargs, + ) + + assert vllm_kwargs["additional_config"]["existing"] == "value" + vllm_config = SimpleNamespace( + additional_config=vllm_kwargs["additional_config"] + ) + assert refit_cache_loader_routes_enabled(vllm_config) is enabled + + @pytest.mark.parametrize( ("tp", "pp", "ep", "expected"), [ diff --git a/tests/unit/reference_configs/eval.yaml b/tests/unit/reference_configs/eval.yaml index abe20f4d74d..f953bd53af1 100644 --- a/tests/unit/reference_configs/eval.yaml +++ b/tests/unit/reference_configs/eval.yaml @@ -19,6 +19,8 @@ generation: vllm_cfg: async_engine: false precision: "bfloat16" + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false tensor_parallel_size: 1 pipeline_parallel_size: 1 expert_parallel_size: 1 diff --git a/tests/unit/utils/test_packed_tensor.py b/tests/unit/utils/test_packed_tensor.py index 6d321bd32aa..f8ee900e38c 100644 --- a/tests/unit/utils/test_packed_tensor.py +++ b/tests/unit/utils/test_packed_tensor.py @@ -18,11 +18,27 @@ import torch from nemo_rl.utils.packed_tensor import ( + _resolve_target_packed_tensor_size, packed_broadcast_consumer, packed_broadcast_producer, ) +def test_explicit_buffer_size_overrides_dynamic_default(): + with patch( + "nemo_rl.utils.packed_tensor.get_target_packed_tensor_size", + return_value=4096, + ): + assert _resolve_target_packed_tensor_size(1024) == 1024 + assert _resolve_target_packed_tensor_size(None) == 4096 + + +@pytest.mark.parametrize("buffer_size_bytes", [0, -1]) +def test_explicit_buffer_size_must_be_positive(buffer_size_bytes): + with pytest.raises(ValueError, match="buffer_size_bytes must be > 0"): + _resolve_target_packed_tensor_size(buffer_size_bytes) + + class MockCommunicationGroup: """Mock communication group for testing broadcast operations.""" diff --git a/tests/unit/weight_sync/test_weight_synchronizer.py b/tests/unit/weight_sync/test_weight_synchronizer.py index 508dab2bcff..1936d6a20a5 100644 --- a/tests/unit/weight_sync/test_weight_synchronizer.py +++ b/tests/unit/weight_sync/test_weight_synchronizer.py @@ -613,8 +613,10 @@ def test_non_colocated_vllm_returns_collective(self): colocated=False, train_cluster=_mock_cluster(), inference_cluster=_mock_cluster(), + refit_buffer_size_gb=1.5, ) assert isinstance(sync, CollectiveWeightSynchronizer) + assert sync._buffer_size_bytes == int(1.5 * 1024**3) def test_non_colocated_sglang_raises(self): policy = _mock_policy() From 2b80a8530b2ab731329dc0a510ca15f719de9a8a Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sat, 1 Aug 2026 12:33:28 -0700 Subject: [PATCH 29/76] fix(refit): cover standalone algorithm paths Signed-off-by: seonjinn --- .../nemo_gym/distillation_qwen3_0_6b.yaml | 2 ++ examples/nemo_gym/grpo_nanov3.yaml | 2 ++ .../nemo_gym/grpo_qwen3_30ba3b_instruct.yaml | 2 ++ .../grpo_qwen3_30ba3b_thinking_swe1.yaml | 2 ++ .../grpo_qwen3_30ba3b_thinking_swe2.yaml | 2 ++ ...rkplace_assistant_nemotron_nano_v2_9b.yaml | 2 ++ .../stage1_rlvr_convergence_27node_h100.yaml | 2 ++ .../stage2_swe1_convergence_20node_h100.yaml | 2 ++ .../stage3_rlhf_convergence_28node_h100.yaml | 2 ++ .../nemotron-3-super/stage1_rlvr.yaml | 2 ++ .../nemotron-3-super/stage2_swe1.yaml | 2 ++ .../nemotron-3-super/stage2_swe2.yaml | 2 ++ .../nemotron-3-super/stage3_rlhf.yaml | 2 ++ .../nemotron-3-ultra/ifbench_teacher.yaml | 2 ++ examples/nemo_gym/nemotron-3-ultra/mopd.yaml | 2 ++ .../nemotron-3-ultra/reasoning_teacher.yaml | 2 ++ .../nemotron-3-ultra/rlhf_teacher.yaml | 2 ++ .../nemotron-3-ultra/student_rlvr1.yaml | 2 ++ .../nemotron-3-ultra/student_rlvr2.yaml | 2 ++ .../nemotron-3-ultra/swe_teacher.yaml | 2 ++ nemo_rl/algorithms/distillation.py | 12 +++++++++-- nemo_rl/algorithms/grpo.py | 4 ++++ nemo_rl/algorithms/grpo_sync.py | 10 +++++++++- nemo_rl/algorithms/ppo.py | 10 +++++++++- tests/unit/utils/test_config.py | 20 +++++++++++++++++++ 25 files changed, 92 insertions(+), 4 deletions(-) diff --git a/examples/nemo_gym/distillation_qwen3_0_6b.yaml b/examples/nemo_gym/distillation_qwen3_0_6b.yaml index 5ee3647f60d..e79b99c6121 100644 --- a/examples/nemo_gym/distillation_qwen3_0_6b.yaml +++ b/examples/nemo_gym/distillation_qwen3_0_6b.yaml @@ -41,6 +41,8 @@ policy: top_p: 1.0 top_k: null vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false async_engine: true expose_http_server: true tensor_parallel_size: 1 diff --git a/examples/nemo_gym/grpo_nanov3.yaml b/examples/nemo_gym/grpo_nanov3.yaml index cb74c313730..8d6927183e0 100644 --- a/examples/nemo_gym/grpo_nanov3.yaml +++ b/examples/nemo_gym/grpo_nanov3.yaml @@ -244,6 +244,8 @@ policy: - deepseek-r1-reasoning - qwen3-coder-tool vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false # NB: can re-enable prefix cache on vllm >= 0.11.2. # enable_prefix_caching: false async_engine: true diff --git a/examples/nemo_gym/grpo_qwen3_30ba3b_instruct.yaml b/examples/nemo_gym/grpo_qwen3_30ba3b_instruct.yaml index 68ff288e2a4..1bc31e8e587 100644 --- a/examples/nemo_gym/grpo_qwen3_30ba3b_instruct.yaml +++ b/examples/nemo_gym/grpo_qwen3_30ba3b_instruct.yaml @@ -52,6 +52,8 @@ policy: generation: vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false tensor_parallel_size: 4 # This is a very low GPU mem utilization. We GPU OOM in two places: # Refit after train, refit before validation. diff --git a/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe1.yaml b/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe1.yaml index 48f284cb711..5ce59297e6b 100644 --- a/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe1.yaml +++ b/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe1.yaml @@ -96,6 +96,8 @@ policy: port_range_high: 4999 max_new_tokens: ${policy.max_total_sequence_length} vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false enable_prefix_caching: true tensor_parallel_size: 2 gpu_memory_utilization: 0.8 diff --git a/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe2.yaml b/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe2.yaml index 6764b43126d..58977f0499e 100644 --- a/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe2.yaml +++ b/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe2.yaml @@ -95,6 +95,8 @@ policy: port_range_high: 4999 max_new_tokens: ${policy.max_total_sequence_length} vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false enable_prefix_caching: true tensor_parallel_size: 2 gpu_memory_utilization: 0.8 diff --git a/examples/nemo_gym/grpo_workplace_assistant_nemotron_nano_v2_9b.yaml b/examples/nemo_gym/grpo_workplace_assistant_nemotron_nano_v2_9b.yaml index 01194e2b770..276248cfdc6 100644 --- a/examples/nemo_gym/grpo_workplace_assistant_nemotron_nano_v2_9b.yaml +++ b/examples/nemo_gym/grpo_workplace_assistant_nemotron_nano_v2_9b.yaml @@ -225,6 +225,8 @@ policy: stop_token_ids: null stop_strings: null vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} tensor_parallel_size: 1 diff --git a/examples/nemo_gym/nemotron-3-super/small_scale/stage1_rlvr_convergence_27node_h100.yaml b/examples/nemo_gym/nemotron-3-super/small_scale/stage1_rlvr_convergence_27node_h100.yaml index e8074ecee50..701233e74cb 100644 --- a/examples/nemo_gym/nemotron-3-super/small_scale/stage1_rlvr_convergence_27node_h100.yaml +++ b/examples/nemo_gym/nemotron-3-super/small_scale/stage1_rlvr_convergence_27node_h100.yaml @@ -7,6 +7,8 @@ policy: resources: num_nodes: 12 vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false max_num_seqs: 16 gpu_memory_utilization: 0.70 diff --git a/examples/nemo_gym/nemotron-3-super/small_scale/stage2_swe1_convergence_20node_h100.yaml b/examples/nemo_gym/nemotron-3-super/small_scale/stage2_swe1_convergence_20node_h100.yaml index 809ecea1723..b71c311ae1f 100644 --- a/examples/nemo_gym/nemotron-3-super/small_scale/stage2_swe1_convergence_20node_h100.yaml +++ b/examples/nemo_gym/nemotron-3-super/small_scale/stage2_swe1_convergence_20node_h100.yaml @@ -10,6 +10,8 @@ policy: resources: num_nodes: 4 # 4 TP8 H100 gen replicas for the 1024-trajectory step vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false max_num_seqs: 16 gpu_memory_utilization: 0.75 diff --git a/examples/nemo_gym/nemotron-3-super/small_scale/stage3_rlhf_convergence_28node_h100.yaml b/examples/nemo_gym/nemotron-3-super/small_scale/stage3_rlhf_convergence_28node_h100.yaml index 0bb802a959b..937c50ddf19 100644 --- a/examples/nemo_gym/nemotron-3-super/small_scale/stage3_rlhf_convergence_28node_h100.yaml +++ b/examples/nemo_gym/nemotron-3-super/small_scale/stage3_rlhf_convergence_28node_h100.yaml @@ -7,6 +7,8 @@ policy: resources: num_nodes: 8 # 8 TP8 H100 gen replicas for the 2048-trajectory step vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false max_num_seqs: 16 gpu_memory_utilization: 0.75 diff --git a/examples/nemo_gym/nemotron-3-super/stage1_rlvr.yaml b/examples/nemo_gym/nemotron-3-super/stage1_rlvr.yaml index f9f48694aa8..69cffdf2e12 100644 --- a/examples/nemo_gym/nemotron-3-super/stage1_rlvr.yaml +++ b/examples/nemo_gym/nemotron-3-super/stage1_rlvr.yaml @@ -225,6 +225,8 @@ policy: stop_token_ids: null stop_strings: null vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-super/stage2_swe1.yaml b/examples/nemo_gym/nemotron-3-super/stage2_swe1.yaml index 941b6327021..d790d4282b9 100644 --- a/examples/nemo_gym/nemotron-3-super/stage2_swe1.yaml +++ b/examples/nemo_gym/nemotron-3-super/stage2_swe1.yaml @@ -225,6 +225,8 @@ policy: stop_token_ids: null stop_strings: null vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false enable_prefix_caching: true async_engine: true precision: ${policy.precision} diff --git a/examples/nemo_gym/nemotron-3-super/stage2_swe2.yaml b/examples/nemo_gym/nemotron-3-super/stage2_swe2.yaml index e714c93b852..11296416752 100644 --- a/examples/nemo_gym/nemotron-3-super/stage2_swe2.yaml +++ b/examples/nemo_gym/nemotron-3-super/stage2_swe2.yaml @@ -218,6 +218,8 @@ policy: stop_token_ids: null stop_strings: null vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false enable_prefix_caching: true async_engine: true precision: ${policy.precision} diff --git a/examples/nemo_gym/nemotron-3-super/stage3_rlhf.yaml b/examples/nemo_gym/nemotron-3-super/stage3_rlhf.yaml index b4a257b4912..b96a322faf5 100644 --- a/examples/nemo_gym/nemotron-3-super/stage3_rlhf.yaml +++ b/examples/nemo_gym/nemotron-3-super/stage3_rlhf.yaml @@ -225,6 +225,8 @@ policy: stop_token_ids: null stop_strings: null vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-ultra/ifbench_teacher.yaml b/examples/nemo_gym/nemotron-3-ultra/ifbench_teacher.yaml index af0e7710f45..adf2e52b326 100644 --- a/examples/nemo_gym/nemotron-3-ultra/ifbench_teacher.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/ifbench_teacher.yaml @@ -307,6 +307,8 @@ policy: # Same memory as EP=1 but uses all-to-all for expert routing. # EP > TP blocked by https://github.com/NVIDIA-NeMo/RL/issues/1101. vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-ultra/mopd.yaml b/examples/nemo_gym/nemotron-3-ultra/mopd.yaml index b7bf33e98e6..d1f364e9345 100644 --- a/examples/nemo_gym/nemotron-3-ultra/mopd.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/mopd.yaml @@ -328,6 +328,8 @@ policy: # Same memory as EP=1 but uses all-to-all for expert routing. # EP > TP blocked by https://github.com/NVIDIA-NeMo/RL/issues/1101. vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-ultra/reasoning_teacher.yaml b/examples/nemo_gym/nemotron-3-ultra/reasoning_teacher.yaml index ef705d41d8b..82a9134e38d 100644 --- a/examples/nemo_gym/nemotron-3-ultra/reasoning_teacher.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/reasoning_teacher.yaml @@ -310,6 +310,8 @@ policy: # Same memory as EP=1 but uses all-to-all for expert routing. # EP > TP blocked by https://github.com/NVIDIA-NeMo/RL/issues/1101. vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-ultra/rlhf_teacher.yaml b/examples/nemo_gym/nemotron-3-ultra/rlhf_teacher.yaml index e7b639baeec..da35ca80d25 100644 --- a/examples/nemo_gym/nemotron-3-ultra/rlhf_teacher.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/rlhf_teacher.yaml @@ -308,6 +308,8 @@ policy: # Same memory as EP=1 but uses all-to-all for expert routing. # EP > TP blocked by https://github.com/NVIDIA-NeMo/RL/issues/1101. vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-ultra/student_rlvr1.yaml b/examples/nemo_gym/nemotron-3-ultra/student_rlvr1.yaml index 883f81752d7..c1fc54ed486 100644 --- a/examples/nemo_gym/nemotron-3-ultra/student_rlvr1.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/student_rlvr1.yaml @@ -304,6 +304,8 @@ policy: # Same memory as EP=1 but uses all-to-all for expert routing. # EP > TP blocked by https://github.com/NVIDIA-NeMo/RL/issues/1101. vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-ultra/student_rlvr2.yaml b/examples/nemo_gym/nemotron-3-ultra/student_rlvr2.yaml index e8481d53a08..796a43c1212 100644 --- a/examples/nemo_gym/nemotron-3-ultra/student_rlvr2.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/student_rlvr2.yaml @@ -305,6 +305,8 @@ policy: # Same memory as EP=1 but uses all-to-all for expert routing. # EP > TP blocked by https://github.com/NVIDIA-NeMo/RL/issues/1101. vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-ultra/swe_teacher.yaml b/examples/nemo_gym/nemotron-3-ultra/swe_teacher.yaml index 8b4cef7d99e..dbd667d351b 100644 --- a/examples/nemo_gym/nemotron-3-ultra/swe_teacher.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/swe_teacher.yaml @@ -324,6 +324,8 @@ policy: stop_token_ids: null stop_strings: null vllm_cfg: + refit_batched_moe_shuffle: true + refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/nemo_rl/algorithms/distillation.py b/nemo_rl/algorithms/distillation.py index faebc9bb58d..754b1d0a11a 100644 --- a/nemo_rl/algorithms/distillation.py +++ b/nemo_rl/algorithms/distillation.py @@ -724,6 +724,7 @@ def distillation_train( val_at_start = master_config.distillation.val_at_start val_at_end = master_config.distillation.val_at_end colocated_inference = master_config.policy["generation"]["colocated"]["enabled"] + refit_buffer_size_gb = master_config.policy.get("refit_buffer_size_gb") max_epochs = ( master_config.distillation.max_num_epochs ) # max number of epochs to train for @@ -736,7 +737,10 @@ def distillation_train( print("\n🔍 Running initial validation...", flush=True) if NEED_REFIT and POLICY_GENERATION_STALE: refit_policy_generation( - student_policy, student_generation, colocated_inference + student_policy, + student_generation, + colocated_inference, + _refit_buffer_size_gb=refit_buffer_size_gb, ) POLICY_GENERATION_STALE = False else: @@ -797,6 +801,7 @@ def distillation_train( student_policy, student_generation, colocated_inference, + _refit_buffer_size_gb=refit_buffer_size_gb, timer=timer, ) POLICY_GENERATION_STALE = False @@ -939,7 +944,10 @@ def distillation_train( ): if NEED_REFIT and POLICY_GENERATION_STALE: refit_policy_generation( - student_policy, student_generation, colocated_inference + student_policy, + student_generation, + colocated_inference, + _refit_buffer_size_gb=refit_buffer_size_gb, ) POLICY_GENERATION_STALE = False else: diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index ee31c906201..e58a0296fad 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -3929,6 +3929,7 @@ def async_grpo_train( val_at_start = master_config.grpo["val_at_start"] val_at_end = master_config.grpo["val_at_end"] colocated_inference = master_config.policy["generation"]["colocated"]["enabled"] + refit_buffer_size_gb = master_config.policy.get("refit_buffer_size_gb") # Initialize advantage estimator adv_estimator = _create_advantage_estimator(master_config) @@ -4081,6 +4082,7 @@ def async_grpo_train( policy, policy_generation, colocated_inference, + _refit_buffer_size_gb=refit_buffer_size_gb, ) print("✅ Policy generation refit completed successfully", flush=True) POLICY_GENERATION_STALE = False @@ -4608,6 +4610,7 @@ def async_grpo_train( policy, policy_generation, colocated_inference, + _refit_buffer_size_gb=refit_buffer_size_gb, ) POLICY_GENERATION_STALE = False @@ -4641,6 +4644,7 @@ def async_grpo_train( policy, policy_generation, colocated_inference, + _refit_buffer_size_gb=refit_buffer_size_gb, ) POLICY_GENERATION_STALE = False else: diff --git a/nemo_rl/algorithms/grpo_sync.py b/nemo_rl/algorithms/grpo_sync.py index e076bccbc6d..1f2ed37c437 100644 --- a/nemo_rl/algorithms/grpo_sync.py +++ b/nemo_rl/algorithms/grpo_sync.py @@ -445,6 +445,7 @@ def grpo_train_sync( val_period = master_config.grpo["val_period"] val_start_at = master_config.grpo["val_start_at"] colocated_inference = master_config.policy["generation"]["colocated"]["enabled"] + refit_buffer_size_gb = master_config.policy.get("refit_buffer_size_gb") # ── Data-plane setup (mandatory in the sync trainer) ─────────────── # Sync trainer requires a TQ-mediated policy. The TQPolicy actor @@ -509,7 +510,12 @@ def grpo_train_sync( memory_tracker.snapshot_start_of_stage("Initial validation", dir()) if NEED_REFIT and POLICY_GENERATION_STALE: - refit_policy_generation(policy, policy_generation, colocated_inference) + refit_policy_generation( + policy, + policy_generation, + colocated_inference, + _refit_buffer_size_gb=refit_buffer_size_gb, + ) POLICY_GENERATION_STALE = False else: policy_generation.prepare_for_generation() @@ -622,6 +628,7 @@ def grpo_train_sync( policy, policy_generation, colocated_inference, + _refit_buffer_size_gb=refit_buffer_size_gb, timer=timer, kv_scales=kv_scales_cache if sync_kv_scales else None, ) @@ -1002,6 +1009,7 @@ def grpo_train_sync( policy, policy_generation, colocated_inference, + _refit_buffer_size_gb=refit_buffer_size_gb, kv_scales=kv_scales_cache if sync_kv_scales else None, ) POLICY_GENERATION_STALE = False diff --git a/nemo_rl/algorithms/ppo.py b/nemo_rl/algorithms/ppo.py index a1adc4335f7..ae4177bf6cd 100644 --- a/nemo_rl/algorithms/ppo.py +++ b/nemo_rl/algorithms/ppo.py @@ -956,6 +956,7 @@ def ppo_train( val_at_end = master_config.ppo["val_at_end"] val_period = master_config.ppo["val_period"] colocated_inference = master_config.policy["generation"]["colocated"]["enabled"] + refit_buffer_size_gb = master_config.policy.get("refit_buffer_size_gb") # Initialize advantage estimator adv_estimator = _create_advantage_estimator(master_config) @@ -966,7 +967,12 @@ def ppo_train( memory_tracker.snapshot_start_of_stage("Initial validation", dir()) if NEED_REFIT and POLICY_GENERATION_STALE: - refit_policy_generation(policy, policy_generation, colocated_inference) + refit_policy_generation( + policy, + policy_generation, + colocated_inference, + _refit_buffer_size_gb=refit_buffer_size_gb, + ) POLICY_GENERATION_STALE = False else: policy_generation.prepare_for_generation() @@ -1057,6 +1063,7 @@ def ppo_train( policy, policy_generation, colocated_inference, + _refit_buffer_size_gb=refit_buffer_size_gb, timer=timer, kv_scales=kv_scales_cache if sync_kv_scales else None, ) @@ -1359,6 +1366,7 @@ def ppo_train( policy, policy_generation, colocated_inference, + _refit_buffer_size_gb=refit_buffer_size_gb, kv_scales=kv_scales_cache if sync_kv_scales else None, ) POLICY_GENERATION_STALE = False diff --git a/tests/unit/utils/test_config.py b/tests/unit/utils/test_config.py index 9d82e89b308..adfeba5d162 100644 --- a/tests/unit/utils/test_config.py +++ b/tests/unit/utils/test_config.py @@ -20,6 +20,14 @@ from nemo_rl.utils.config import load_config, register_omegaconf_resolvers REPO_ROOT = Path(__file__).resolve().parents[3] +NEMO_GYM_VLLM_CONFIG_PATHS = [ + config_path.relative_to(REPO_ROOT) + for config_path in sorted((REPO_ROOT / "examples/nemo_gym").rglob("*.yaml")) + if OmegaConf.select( + OmegaConf.load(config_path), "policy.generation.vllm_cfg" + ) + is not None +] ULTRA_CONFIG_PATHS = [ "examples/nemo_gym/nemotron-3-ultra/student_rlvr1.yaml", "examples/nemo_gym/nemotron-3-ultra/student_rlvr2.yaml", @@ -232,6 +240,18 @@ def test_add_resolver(): assert config.value == 5 +@pytest.mark.parametrize("config_path", NEMO_GYM_VLLM_CONFIG_PATHS) +def test_nemo_gym_vllm_configs_define_refit_defaults(config_path): + """Ensure standalone NeMo-Gym vLLM configs set required refit defaults.""" + config = OmegaConf.load(REPO_ROOT / config_path) + vllm_config = config.policy.generation.vllm_cfg + + assert "refit_batched_moe_shuffle" in vllm_config + assert vllm_config.refit_batched_moe_shuffle is True + assert "refit_cache_loader_routes" in vllm_config + assert vllm_config.refit_cache_loader_routes is False + + @pytest.mark.parametrize("config_path", ULTRA_CONFIG_PATHS) def test_ultra_configs_satisfy_current_grpo_contract(config_path): """Ensure Ultra configs compose with all fields required by current GRPO.""" From 0b6b5cc1acb2f621fbfeae259250266274e0d093 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sat, 1 Aug 2026 13:57:37 -0700 Subject: [PATCH 30/76] fix(vllm): scope refit optimization config Signed-off-by: seonjinn --- examples/configs/distillation_math.yaml | 2 -- examples/configs/evals/eval.yaml | 2 -- examples/configs/evals/mmau.yaml | 2 -- examples/configs/grpo_math_1B.yaml | 2 -- examples/configs/ppo_math_1B.yaml | 2 -- .../nemo_gym/distillation_qwen3_0_6b.yaml | 2 -- examples/nemo_gym/grpo_nanov3.yaml | 2 -- .../nemo_gym/grpo_qwen3_30ba3b_instruct.yaml | 2 -- .../grpo_qwen3_30ba3b_thinking_swe1.yaml | 2 -- .../grpo_qwen3_30ba3b_thinking_swe2.yaml | 2 -- ...rkplace_assistant_nemotron_nano_v2_9b.yaml | 2 -- .../stage1_rlvr_convergence_27node_h100.yaml | 2 -- .../stage2_swe1_convergence_20node_h100.yaml | 2 -- .../stage3_rlhf_convergence_28node_h100.yaml | 2 -- .../nemotron-3-super/stage1_rlvr.yaml | 2 -- .../nemotron-3-super/stage2_swe1.yaml | 2 -- .../nemotron-3-super/stage2_swe2.yaml | 2 -- .../nemotron-3-super/stage3_rlhf.yaml | 2 -- .../nemotron-3-ultra/ifbench_teacher.yaml | 2 -- examples/nemo_gym/nemotron-3-ultra/mopd.yaml | 2 -- .../nemotron-3-ultra/reasoning_teacher.yaml | 2 -- .../nemotron-3-ultra/rlhf_teacher.yaml | 2 -- .../nemotron-3-ultra/student_rlvr1.yaml | 2 -- .../nemotron-3-ultra/student_rlvr2.yaml | 2 -- .../nemotron-3-ultra/swe_teacher.yaml | 2 -- nemo_rl/models/generation/vllm/config.py | 4 +-- .../generation/vllm/quantization/fp8.py | 10 +++----- .../models/generation/vllm/vllm_backend.py | 7 +++--- .../models/generation/vllm/worker_utils.py | 7 +++--- .../generation/test_vllm_fp8_quantization.py | 25 ++++++++++++++++++- .../generation/test_vllm_worker_helpers.py | 13 +++++++--- tests/unit/utils/test_config.py | 20 --------------- 32 files changed, 46 insertions(+), 90 deletions(-) diff --git a/examples/configs/distillation_math.yaml b/examples/configs/distillation_math.yaml index b8bced51dc0..11ca9d09072 100644 --- a/examples/configs/distillation_math.yaml +++ b/examples/configs/distillation_math.yaml @@ -205,8 +205,6 @@ policy: &POLICY_BASE vllm_cfg: async_engine: false precision: ${...precision} - refit_batched_moe_shuffle: true # Batch MoE layout transforms across experts during refit. - refit_cache_loader_routes: false # Replay stable vLLM loader routes across refits. kv_cache_dtype: "auto" tensor_parallel_size: 1 pipeline_parallel_size: 1 diff --git a/examples/configs/evals/eval.yaml b/examples/configs/evals/eval.yaml index 05859cde775..c492ac34cd1 100644 --- a/examples/configs/evals/eval.yaml +++ b/examples/configs/evals/eval.yaml @@ -19,8 +19,6 @@ generation: vllm_cfg: async_engine: false precision: "bfloat16" - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false tensor_parallel_size: 1 pipeline_parallel_size: 1 expert_parallel_size: 1 diff --git a/examples/configs/evals/mmau.yaml b/examples/configs/evals/mmau.yaml index fafe2a79155..e12c3ea0aec 100644 --- a/examples/configs/evals/mmau.yaml +++ b/examples/configs/evals/mmau.yaml @@ -18,8 +18,6 @@ generation: vllm_cfg: async_engine: false precision: "bfloat16" - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false tensor_parallel_size: 1 pipeline_parallel_size: 1 expert_parallel_size: 1 diff --git a/examples/configs/grpo_math_1B.yaml b/examples/configs/grpo_math_1B.yaml index bfdf2fea083..ee464e2fec2 100644 --- a/examples/configs/grpo_math_1B.yaml +++ b/examples/configs/grpo_math_1B.yaml @@ -388,8 +388,6 @@ policy: precision: ${policy.precision} # MXFP8 + Megatron only: quantize on trainer and stream E4M3 values plus scales. refit_prequantize: false - refit_batched_moe_shuffle: true # Batch MoE layout transforms across experts during refit. - refit_cache_loader_routes: false # Replay stable vLLM loader routes across refits. kv_cache_dtype: "auto" tensor_parallel_size: 1 pipeline_parallel_size: 1 diff --git a/examples/configs/ppo_math_1B.yaml b/examples/configs/ppo_math_1B.yaml index 2dc9a346b08..ceed0e3c4ef 100644 --- a/examples/configs/ppo_math_1B.yaml +++ b/examples/configs/ppo_math_1B.yaml @@ -247,8 +247,6 @@ policy: vllm_cfg: async_engine: false precision: ${policy.precision} - refit_batched_moe_shuffle: true # Batch MoE layout transforms across experts during refit. - refit_cache_loader_routes: false # Replay stable vLLM loader routes across refits. kv_cache_dtype: "auto" tensor_parallel_size: 1 pipeline_parallel_size: 1 diff --git a/examples/nemo_gym/distillation_qwen3_0_6b.yaml b/examples/nemo_gym/distillation_qwen3_0_6b.yaml index e79b99c6121..5ee3647f60d 100644 --- a/examples/nemo_gym/distillation_qwen3_0_6b.yaml +++ b/examples/nemo_gym/distillation_qwen3_0_6b.yaml @@ -41,8 +41,6 @@ policy: top_p: 1.0 top_k: null vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false async_engine: true expose_http_server: true tensor_parallel_size: 1 diff --git a/examples/nemo_gym/grpo_nanov3.yaml b/examples/nemo_gym/grpo_nanov3.yaml index 8d6927183e0..cb74c313730 100644 --- a/examples/nemo_gym/grpo_nanov3.yaml +++ b/examples/nemo_gym/grpo_nanov3.yaml @@ -244,8 +244,6 @@ policy: - deepseek-r1-reasoning - qwen3-coder-tool vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false # NB: can re-enable prefix cache on vllm >= 0.11.2. # enable_prefix_caching: false async_engine: true diff --git a/examples/nemo_gym/grpo_qwen3_30ba3b_instruct.yaml b/examples/nemo_gym/grpo_qwen3_30ba3b_instruct.yaml index 1bc31e8e587..68ff288e2a4 100644 --- a/examples/nemo_gym/grpo_qwen3_30ba3b_instruct.yaml +++ b/examples/nemo_gym/grpo_qwen3_30ba3b_instruct.yaml @@ -52,8 +52,6 @@ policy: generation: vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false tensor_parallel_size: 4 # This is a very low GPU mem utilization. We GPU OOM in two places: # Refit after train, refit before validation. diff --git a/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe1.yaml b/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe1.yaml index 5ce59297e6b..48f284cb711 100644 --- a/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe1.yaml +++ b/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe1.yaml @@ -96,8 +96,6 @@ policy: port_range_high: 4999 max_new_tokens: ${policy.max_total_sequence_length} vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false enable_prefix_caching: true tensor_parallel_size: 2 gpu_memory_utilization: 0.8 diff --git a/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe2.yaml b/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe2.yaml index 58977f0499e..6764b43126d 100644 --- a/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe2.yaml +++ b/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe2.yaml @@ -95,8 +95,6 @@ policy: port_range_high: 4999 max_new_tokens: ${policy.max_total_sequence_length} vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false enable_prefix_caching: true tensor_parallel_size: 2 gpu_memory_utilization: 0.8 diff --git a/examples/nemo_gym/grpo_workplace_assistant_nemotron_nano_v2_9b.yaml b/examples/nemo_gym/grpo_workplace_assistant_nemotron_nano_v2_9b.yaml index 276248cfdc6..01194e2b770 100644 --- a/examples/nemo_gym/grpo_workplace_assistant_nemotron_nano_v2_9b.yaml +++ b/examples/nemo_gym/grpo_workplace_assistant_nemotron_nano_v2_9b.yaml @@ -225,8 +225,6 @@ policy: stop_token_ids: null stop_strings: null vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} tensor_parallel_size: 1 diff --git a/examples/nemo_gym/nemotron-3-super/small_scale/stage1_rlvr_convergence_27node_h100.yaml b/examples/nemo_gym/nemotron-3-super/small_scale/stage1_rlvr_convergence_27node_h100.yaml index 701233e74cb..e8074ecee50 100644 --- a/examples/nemo_gym/nemotron-3-super/small_scale/stage1_rlvr_convergence_27node_h100.yaml +++ b/examples/nemo_gym/nemotron-3-super/small_scale/stage1_rlvr_convergence_27node_h100.yaml @@ -7,8 +7,6 @@ policy: resources: num_nodes: 12 vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false max_num_seqs: 16 gpu_memory_utilization: 0.70 diff --git a/examples/nemo_gym/nemotron-3-super/small_scale/stage2_swe1_convergence_20node_h100.yaml b/examples/nemo_gym/nemotron-3-super/small_scale/stage2_swe1_convergence_20node_h100.yaml index b71c311ae1f..809ecea1723 100644 --- a/examples/nemo_gym/nemotron-3-super/small_scale/stage2_swe1_convergence_20node_h100.yaml +++ b/examples/nemo_gym/nemotron-3-super/small_scale/stage2_swe1_convergence_20node_h100.yaml @@ -10,8 +10,6 @@ policy: resources: num_nodes: 4 # 4 TP8 H100 gen replicas for the 1024-trajectory step vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false max_num_seqs: 16 gpu_memory_utilization: 0.75 diff --git a/examples/nemo_gym/nemotron-3-super/small_scale/stage3_rlhf_convergence_28node_h100.yaml b/examples/nemo_gym/nemotron-3-super/small_scale/stage3_rlhf_convergence_28node_h100.yaml index 937c50ddf19..0bb802a959b 100644 --- a/examples/nemo_gym/nemotron-3-super/small_scale/stage3_rlhf_convergence_28node_h100.yaml +++ b/examples/nemo_gym/nemotron-3-super/small_scale/stage3_rlhf_convergence_28node_h100.yaml @@ -7,8 +7,6 @@ policy: resources: num_nodes: 8 # 8 TP8 H100 gen replicas for the 2048-trajectory step vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false max_num_seqs: 16 gpu_memory_utilization: 0.75 diff --git a/examples/nemo_gym/nemotron-3-super/stage1_rlvr.yaml b/examples/nemo_gym/nemotron-3-super/stage1_rlvr.yaml index 69cffdf2e12..f9f48694aa8 100644 --- a/examples/nemo_gym/nemotron-3-super/stage1_rlvr.yaml +++ b/examples/nemo_gym/nemotron-3-super/stage1_rlvr.yaml @@ -225,8 +225,6 @@ policy: stop_token_ids: null stop_strings: null vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-super/stage2_swe1.yaml b/examples/nemo_gym/nemotron-3-super/stage2_swe1.yaml index d790d4282b9..941b6327021 100644 --- a/examples/nemo_gym/nemotron-3-super/stage2_swe1.yaml +++ b/examples/nemo_gym/nemotron-3-super/stage2_swe1.yaml @@ -225,8 +225,6 @@ policy: stop_token_ids: null stop_strings: null vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false enable_prefix_caching: true async_engine: true precision: ${policy.precision} diff --git a/examples/nemo_gym/nemotron-3-super/stage2_swe2.yaml b/examples/nemo_gym/nemotron-3-super/stage2_swe2.yaml index 11296416752..e714c93b852 100644 --- a/examples/nemo_gym/nemotron-3-super/stage2_swe2.yaml +++ b/examples/nemo_gym/nemotron-3-super/stage2_swe2.yaml @@ -218,8 +218,6 @@ policy: stop_token_ids: null stop_strings: null vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false enable_prefix_caching: true async_engine: true precision: ${policy.precision} diff --git a/examples/nemo_gym/nemotron-3-super/stage3_rlhf.yaml b/examples/nemo_gym/nemotron-3-super/stage3_rlhf.yaml index b96a322faf5..b4a257b4912 100644 --- a/examples/nemo_gym/nemotron-3-super/stage3_rlhf.yaml +++ b/examples/nemo_gym/nemotron-3-super/stage3_rlhf.yaml @@ -225,8 +225,6 @@ policy: stop_token_ids: null stop_strings: null vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-ultra/ifbench_teacher.yaml b/examples/nemo_gym/nemotron-3-ultra/ifbench_teacher.yaml index adf2e52b326..af0e7710f45 100644 --- a/examples/nemo_gym/nemotron-3-ultra/ifbench_teacher.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/ifbench_teacher.yaml @@ -307,8 +307,6 @@ policy: # Same memory as EP=1 but uses all-to-all for expert routing. # EP > TP blocked by https://github.com/NVIDIA-NeMo/RL/issues/1101. vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-ultra/mopd.yaml b/examples/nemo_gym/nemotron-3-ultra/mopd.yaml index d1f364e9345..b7bf33e98e6 100644 --- a/examples/nemo_gym/nemotron-3-ultra/mopd.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/mopd.yaml @@ -328,8 +328,6 @@ policy: # Same memory as EP=1 but uses all-to-all for expert routing. # EP > TP blocked by https://github.com/NVIDIA-NeMo/RL/issues/1101. vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-ultra/reasoning_teacher.yaml b/examples/nemo_gym/nemotron-3-ultra/reasoning_teacher.yaml index 82a9134e38d..ef705d41d8b 100644 --- a/examples/nemo_gym/nemotron-3-ultra/reasoning_teacher.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/reasoning_teacher.yaml @@ -310,8 +310,6 @@ policy: # Same memory as EP=1 but uses all-to-all for expert routing. # EP > TP blocked by https://github.com/NVIDIA-NeMo/RL/issues/1101. vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-ultra/rlhf_teacher.yaml b/examples/nemo_gym/nemotron-3-ultra/rlhf_teacher.yaml index da35ca80d25..e7b639baeec 100644 --- a/examples/nemo_gym/nemotron-3-ultra/rlhf_teacher.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/rlhf_teacher.yaml @@ -308,8 +308,6 @@ policy: # Same memory as EP=1 but uses all-to-all for expert routing. # EP > TP blocked by https://github.com/NVIDIA-NeMo/RL/issues/1101. vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-ultra/student_rlvr1.yaml b/examples/nemo_gym/nemotron-3-ultra/student_rlvr1.yaml index c1fc54ed486..883f81752d7 100644 --- a/examples/nemo_gym/nemotron-3-ultra/student_rlvr1.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/student_rlvr1.yaml @@ -304,8 +304,6 @@ policy: # Same memory as EP=1 but uses all-to-all for expert routing. # EP > TP blocked by https://github.com/NVIDIA-NeMo/RL/issues/1101. vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-ultra/student_rlvr2.yaml b/examples/nemo_gym/nemotron-3-ultra/student_rlvr2.yaml index 796a43c1212..e8481d53a08 100644 --- a/examples/nemo_gym/nemotron-3-ultra/student_rlvr2.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/student_rlvr2.yaml @@ -305,8 +305,6 @@ policy: # Same memory as EP=1 but uses all-to-all for expert routing. # EP > TP blocked by https://github.com/NVIDIA-NeMo/RL/issues/1101. vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/examples/nemo_gym/nemotron-3-ultra/swe_teacher.yaml b/examples/nemo_gym/nemotron-3-ultra/swe_teacher.yaml index dbd667d351b..8b4cef7d99e 100644 --- a/examples/nemo_gym/nemotron-3-ultra/swe_teacher.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/swe_teacher.yaml @@ -324,8 +324,6 @@ policy: stop_token_ids: null stop_strings: null vllm_cfg: - refit_batched_moe_shuffle: true - refit_cache_loader_routes: false async_engine: true precision: ${policy.precision} kv_cache_dtype: "auto" diff --git a/nemo_rl/models/generation/vllm/config.py b/nemo_rl/models/generation/vllm/config.py index dea33957110..4a9c32238f4 100644 --- a/nemo_rl/models/generation/vllm/config.py +++ b/nemo_rl/models/generation/vllm/config.py @@ -42,9 +42,9 @@ class VllmSpecificArgs(TypedDict): # Megatron policy backend. refit_prequantize: NotRequired[bool] # Batch MoE weight-layout transforms across experts during MXFP8 refit. - refit_batched_moe_shuffle: bool + refit_batched_moe_shuffle: NotRequired[bool] # Cache and replay stable vLLM weight-loader routes across refits. - refit_cache_loader_routes: bool + refit_cache_loader_routes: NotRequired[bool] kv_cache_dtype: Literal["auto", "fp8", "fp8_e4m3"] enforce_eager: NotRequired[bool] enable_return_routed_experts: NotRequired[bool] diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index 937b2c4cf29..940e73b7935 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -114,9 +114,7 @@ def install_fp8_config(config: dict[str, Any] | None) -> None: global_fp8_config = FP8Config(**config) -def _patch_ray_executor_v2_worker( - ray_executor_v2: Any, fp8_config: FP8Config -) -> None: +def _patch_ray_executor_v2_worker(ray_executor_v2: Any, fp8_config: FP8Config) -> None: """Install FP8 patches inside RayExecutorV2 workers before model loading.""" original_ray_worker_proc = ray_executor_v2.RayWorkerProc if getattr(original_ray_worker_proc, "_nrl_fp8_patched", False): @@ -284,7 +282,7 @@ def init_fp8(vllm_cfg, model_name, model_parallel_size): "model_parallel_size": model_parallel_size, "kv_cache_dtype": kv_cache_dtype, "use_fp8_weights": use_fp8_weights, - "refit_batched_moe_shuffle": vllm_cfg["refit_batched_moe_shuffle"], + "refit_batched_moe_shuffle": vllm_cfg.get("refit_batched_moe_shuffle", True), } if is_mx: fp8_config_kwargs["is_mx"] = True @@ -563,9 +561,7 @@ def load_weights(weights, model_runner): load_weights_maybe_cached( model, weights_quantized, - cache_loader_routes=refit_cache_loader_routes_enabled( - model_runner.vllm_config - ), + cache_loader_routes=refit_cache_loader_routes_enabled(model_runner.vllm_config), ) diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index e6671fc8b24..c3db3a492b8 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -28,14 +28,14 @@ preinit_nixl_from_vllm_config, resolve_rollout_rank, ) +from nemo_rl.models.generation.vllm.worker_utils import ( + refit_cache_loader_routes_enabled, +) from nemo_rl.models.policy.utils import ( IPCProtocol, calculate_aligned_size, rebuild_cuda_tensor_from_ipc, ) -from nemo_rl.models.generation.vllm.worker_utils import ( - refit_cache_loader_routes_enabled, -) from nemo_rl.utils.nsys import wrap_with_nvtx_name from nemo_rl.utils.packed_tensor import packed_broadcast_consumer from nemo_rl.weight_sync.nccl_reshard_utils import ( @@ -320,6 +320,7 @@ class VllmInternalWorkerExtension: _mtp_drafter_from_disk: bool = False _sparse_delta_applier: Any = None _nrl_named_parameters: dict[str, torch.nn.Parameter] + def _get_named_parameters(self) -> dict[str, torch.nn.Parameter]: params = getattr(self, "_nrl_named_parameters", None) if params is None: diff --git a/nemo_rl/models/generation/vllm/worker_utils.py b/nemo_rl/models/generation/vllm/worker_utils.py index 70cf996e0d1..2c22170e8fd 100644 --- a/nemo_rl/models/generation/vllm/worker_utils.py +++ b/nemo_rl/models/generation/vllm/worker_utils.py @@ -15,7 +15,6 @@ from collections.abc import Mapping from typing import Any - _REFIT_CACHE_LOADER_ROUTES_KEY = "nemo_rl_refit_cache_loader_routes" @@ -24,9 +23,9 @@ def configure_refit_runtime( ) -> None: """Forward NeMo-RL refit options through vLLM's worker config.""" additional_config = dict(vllm_kwargs.get("additional_config") or {}) - additional_config[_REFIT_CACHE_LOADER_ROUTES_KEY] = vllm_cfg[ - "refit_cache_loader_routes" - ] + additional_config[_REFIT_CACHE_LOADER_ROUTES_KEY] = vllm_cfg.get( + "refit_cache_loader_routes", False + ) vllm_kwargs["additional_config"] = additional_config diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 33c1f2e1345..511219eba24 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -12,10 +12,10 @@ # See the License for the specific language governing permissions and # limitations under the License. -import cloudpickle import types from typing import Any +import cloudpickle import pytest import torch @@ -92,6 +92,29 @@ def test_init_fp8_uses_mxfp8_quantization_config(fp8_module, monkeypatch): assert "VLLM_USE_DEEP_GEMM_E8M0" not in fp8.os.environ +def test_init_fp8_defaults_to_batched_moe_shuffle(fp8_module, monkeypatch): + fp8 = fp8_module + monkeypatch.setattr( + fp8.AutoConfig, + "from_pretrained", + lambda *_args, **_kwargs: types.SimpleNamespace(num_hidden_layers=4), + ) + monkeypatch.setattr(fp8, "monkey_patch_vllm_ray_executor", lambda _config: None) + + fp8.init_fp8( + { + "precision": "fp8", + "kv_cache_dtype": "auto", + "async_engine": False, + "is_mx": True, + }, + "dummy-model", + model_parallel_size=1, + ) + + assert fp8.global_fp8_config.refit_batched_moe_shuffle is True + + def test_ray_executor_v2_worker_applies_fp8_patches_before_model_load( fp8_module, monkeypatch ): diff --git a/tests/unit/models/generation/test_vllm_worker_helpers.py b/tests/unit/models/generation/test_vllm_worker_helpers.py index e600ff82f1a..9ef0e6de842 100644 --- a/tests/unit/models/generation/test_vllm_worker_helpers.py +++ b/tests/unit/models/generation/test_vllm_worker_helpers.py @@ -36,12 +36,19 @@ def test_refit_loader_cache_round_trips_through_additional_config(enabled): ) assert vllm_kwargs["additional_config"]["existing"] == "value" - vllm_config = SimpleNamespace( - additional_config=vllm_kwargs["additional_config"] - ) + vllm_config = SimpleNamespace(additional_config=vllm_kwargs["additional_config"]) assert refit_cache_loader_routes_enabled(vllm_config) is enabled +def test_refit_loader_cache_defaults_to_disabled(): + vllm_kwargs = {} + + configure_refit_runtime({}, vllm_kwargs) + + vllm_config = SimpleNamespace(additional_config=vllm_kwargs["additional_config"]) + assert refit_cache_loader_routes_enabled(vllm_config) is False + + @pytest.mark.parametrize( ("tp", "pp", "ep", "expected"), [ diff --git a/tests/unit/utils/test_config.py b/tests/unit/utils/test_config.py index adfeba5d162..9d82e89b308 100644 --- a/tests/unit/utils/test_config.py +++ b/tests/unit/utils/test_config.py @@ -20,14 +20,6 @@ from nemo_rl.utils.config import load_config, register_omegaconf_resolvers REPO_ROOT = Path(__file__).resolve().parents[3] -NEMO_GYM_VLLM_CONFIG_PATHS = [ - config_path.relative_to(REPO_ROOT) - for config_path in sorted((REPO_ROOT / "examples/nemo_gym").rglob("*.yaml")) - if OmegaConf.select( - OmegaConf.load(config_path), "policy.generation.vllm_cfg" - ) - is not None -] ULTRA_CONFIG_PATHS = [ "examples/nemo_gym/nemotron-3-ultra/student_rlvr1.yaml", "examples/nemo_gym/nemotron-3-ultra/student_rlvr2.yaml", @@ -240,18 +232,6 @@ def test_add_resolver(): assert config.value == 5 -@pytest.mark.parametrize("config_path", NEMO_GYM_VLLM_CONFIG_PATHS) -def test_nemo_gym_vllm_configs_define_refit_defaults(config_path): - """Ensure standalone NeMo-Gym vLLM configs set required refit defaults.""" - config = OmegaConf.load(REPO_ROOT / config_path) - vllm_config = config.policy.generation.vllm_cfg - - assert "refit_batched_moe_shuffle" in vllm_config - assert vllm_config.refit_batched_moe_shuffle is True - assert "refit_cache_loader_routes" in vllm_config - assert vllm_config.refit_cache_loader_routes is False - - @pytest.mark.parametrize("config_path", ULTRA_CONFIG_PATHS) def test_ultra_configs_satisfy_current_grpo_contract(config_path): """Ensure Ultra configs compose with all fields required by current GRPO.""" From cd7732c4988d361f43148e137f848a167ef00355 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sat, 1 Aug 2026 14:38:05 -0700 Subject: [PATCH 31/76] fix(refit): allow matching fp8 reshard storage Signed-off-by: seonjinn --- nemo_rl/weight_sync/nccl_reshard_utils.py | 8 -------- tests/unit/weight_sync/test_nccl_reshard_utils.py | 12 +----------- 2 files changed, 1 insertion(+), 19 deletions(-) diff --git a/nemo_rl/weight_sync/nccl_reshard_utils.py b/nemo_rl/weight_sync/nccl_reshard_utils.py index 9bfb4c9dfb6..92a4b53feef 100644 --- a/nemo_rl/weight_sync/nccl_reshard_utils.py +++ b/nemo_rl/weight_sync/nccl_reshard_utils.py @@ -535,7 +535,6 @@ def check_nccl_reshard_refit_support(master_config: dict) -> None: generation = policy.get("generation", {}) or {} megatron_cfg = policy.get("megatron_cfg", {}) or {} dtensor_cfg = policy.get("dtensor_cfg", {}) or {} - policy_precision = policy.get("precision") vllm_cfg = generation.get("vllm_cfg", {}) or {} vllm_kwargs = generation.get("vllm_kwargs", {}) or {} @@ -555,13 +554,6 @@ def check_nccl_reshard_refit_support(master_config: dict) -> None: f"policy.generation.backend must be 'vllm' (got {backend!r})." ) - if policy_precision != "bfloat16": - violations.append( - "policy.precision must be 'bfloat16' for nccl_reshard_refit " - f"(got {policy_precision!r}); the refit byte-copies training storage " - "into the generation model." - ) - if vllm_kwargs.get("enable_eplb"): violations.append( "policy.generation.vllm_kwargs.enable_eplb must be False " diff --git a/tests/unit/weight_sync/test_nccl_reshard_utils.py b/tests/unit/weight_sync/test_nccl_reshard_utils.py index 0f96668211f..b5b5b7b601b 100644 --- a/tests/unit/weight_sync/test_nccl_reshard_utils.py +++ b/tests/unit/weight_sync/test_nccl_reshard_utils.py @@ -115,6 +115,7 @@ def test_check_nccl_reshard_refit_support_rejects_unsupported_refit_modes( def test_check_nccl_reshard_refit_support_accepts_matching_blockwise_fp8() -> None: config = _valid_nccl_reshard_config() + config.policy["precision"] = "fp8" config.policy["generation"]["vllm_cfg"]["precision"] = "fp8" config.policy["megatron_cfg"]["fp8_cfg"] = { "enabled": True, @@ -138,17 +139,6 @@ def test_check_nccl_reshard_refit_support_rejects_disabled_fp8_param_storage() - check_nccl_reshard_refit_support(config) -@pytest.mark.parametrize("precision", ["float16", "float32"]) -def test_check_nccl_reshard_refit_support_rejects_non_bfloat16_policy_precision( - precision: str, -) -> None: - config = _valid_nccl_reshard_config() - config.policy["precision"] = precision - - with pytest.raises(ValueError, match="policy.precision must be 'bfloat16'"): - check_nccl_reshard_refit_support(config) - - # -------------------------------------------------------------------------- # MeshInfo # -------------------------------------------------------------------------- From 7705d22a3138f9787158c6c54f4e4282c24eea0e Mon Sep 17 00:00:00 2001 From: seonjinn Date: Thu, 13 Aug 2026 14:19:44 -0700 Subject: [PATCH 32/76] fix(refit): validate MXFP8 prequant wire format Signed-off-by: seonjinn --- nemo_rl/models/generation/vllm/config.py | 5 + .../generation/vllm/quantization/fp8.py | 16 +- .../policy/workers/megatron_policy_worker.py | 16 +- .../models/generation/test_mxfp8_prequant.py | 39 ----- .../models/generation/test_vllm_config.py | 17 +++ .../generation/test_vllm_fp8_quantization.py | 139 ++++++++++++++++-- .../models/policy/test_megatron_worker.py | 38 ++++- 7 files changed, 205 insertions(+), 65 deletions(-) diff --git a/nemo_rl/models/generation/vllm/config.py b/nemo_rl/models/generation/vllm/config.py index f569ca585e2..7edacbb22a4 100644 --- a/nemo_rl/models/generation/vllm/config.py +++ b/nemo_rl/models/generation/vllm/config.py @@ -186,6 +186,11 @@ def validate_vllm_quantization_config(config: VllmConfig) -> None: "policy.generation.vllm_cfg.refit_prequantize requires " "precision='fp8' and is_mx=true." ) + if refit_prequantize and config.get("refit_transport") == "nccl_reshard": + raise ValueError( + "policy.generation.vllm_cfg.refit_prequantize is not supported with " + "nccl_reshard; that transport owns its weight-format conversion." + ) for field in ("refit_cache_loader_routes",): value = vllm_cfg.get(field) if value is not None and not isinstance(value, bool): diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index 1b7cc26e8fb..6b4ed803b23 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -514,6 +514,8 @@ def _is_fp8_weight(name, model): def load_weights(weights, model_runner): global global_fp8_config + weights = list(weights) + weight_names = {name for name, _tensor in weights} weights_quantized = [] model = model_runner.model @@ -522,8 +524,18 @@ def load_weights(weights, model_runner): weights_quantized.append((k, v)) continue if v.dtype == torch.float8_e4m3fn: - # Already quantized on the trainer (vllm_cfg.refit_prequantize); the - # matching *_scale_from_checkpoint entry arrives as its own weight. + if global_fp8_config.is_mx and not global_fp8_config.refit_prequantize: + raise ValueError( + "MXFP8 E4M3 refit weights require refit_prequantize=true; " + "other FP8 trainer scale layouts are not compatible." + ) + scale_name = k + "_scale_from_checkpoint" + if global_fp8_config.is_mx and scale_name not in weight_names: + raise ValueError( + f"Prequantized MXFP8 weight {k!r} is missing {scale_name!r}." + ) + # Prequantized MXFP8 sends the matching *_scale_from_checkpoint + # entry separately. Non-MXFP8 blockwise FP8 sends *_scale_inv. weights_quantized.append([k, v]) continue # Cast the weight into fp8 and its scale factor diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 1dfbb2e2f4b..2bb476a0b78 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -1973,6 +1973,12 @@ def enable_refit_prequantize(self, param_names: list[str]) -> dict[str, Any]: Updated refit metadata: the listed params become float8_e4m3fn and each gains a *_scale_from_checkpoint uint8 entry. """ + if self._is_fp8_export(): + raise ValueError( + "vllm_cfg.refit_prequantize requires BF16 trainer-exported weights; " + "Megatron blockwise FP8 parameter storage uses a different scale layout." + ) + self._refit_prequant_names = set(param_names) refit_param_info_hf = {} @@ -1983,12 +1989,14 @@ def enable_refit_prequantize(self, param_names: list[str]) -> dict[str, Any]: def _maybe_prequantize_param( self, name: str, tensor: torch.Tensor ) -> Iterator[tuple[str, torch.Tensor]]: - if ( - name not in self._refit_prequant_names - or tensor.dtype == torch.float8_e4m3fn - ): + if name not in self._refit_prequant_names: yield name, tensor return + if tensor.dtype == torch.float8_e4m3fn: + raise ValueError( + "vllm_cfg.refit_prequantize requires BF16 trainer-exported weights; " + f"{name} is already stored as E4M3 with a non-MXFP8 scale layout." + ) # Deferred: pulls in the heavy nemo_rl...generation.vllm package init, # which trainer workers only need when prequantized refit is enabled. diff --git a/tests/unit/models/generation/test_mxfp8_prequant.py b/tests/unit/models/generation/test_mxfp8_prequant.py index 4d3f848c825..8eaf5b2564b 100644 --- a/tests/unit/models/generation/test_mxfp8_prequant.py +++ b/tests/unit/models/generation/test_mxfp8_prequant.py @@ -172,42 +172,3 @@ def rand_bytes(*shape): assert got.shape == want.shape, name assert got.dtype == want.dtype, name assert torch.equal(got.view(torch.uint8), want.view(torch.uint8)), name - - -def test_mxfp8_shuffle_verification_runs_once_per_layer(monkeypatch): - fp8 = pytest.importorskip("nemo_rl.models.generation.vllm.quantization.fp8") - - class Layer: - pass - - monkeypatch.setenv("NRL_MXFP8_SHUFFLE_VERIFY", "1") - fp8.mxfp8_shuffle_verified_layers.clear() - tensors = ( - torch.arange(8, dtype=torch.uint8), - torch.arange(8, dtype=torch.uint8), - torch.arange(8, dtype=torch.uint8), - torch.arange(8, dtype=torch.uint8), - ) - calls = [] - - def reference(*args): - calls.append(args[0]) - return tensors - - monkeypatch.setattr(fp8, "_shuffle_mxfp8_moe_per_expert", reference) - layers = [Layer(), Layer()] - for layer in layers: - for _ in range(2): - fp8._verify_mxfp8_moe_shuffle( - layer, - tensors[0], - tensors[1], - tensors[2], - tensors[3], - False, - 128, - tensors, - ) - - assert len(calls) == len(layers) - assert all(layer in fp8.mxfp8_shuffle_verified_layers for layer in layers) diff --git a/tests/unit/models/generation/test_vllm_config.py b/tests/unit/models/generation/test_vllm_config.py index 1f79a105ed6..e4043878d93 100644 --- a/tests/unit/models/generation/test_vllm_config.py +++ b/tests/unit/models/generation/test_vllm_config.py @@ -90,6 +90,23 @@ def test_refit_prequantize_accepts_mxfp8() -> None: validate_vllm_quantization_config(generation_config) +def test_refit_prequantize_rejects_nccl_reshard() -> None: + generation_config = cast( + VllmConfig, + { + "refit_transport": "nccl_reshard", + "vllm_cfg": { + "precision": "fp8", + "is_mx": True, + "refit_prequantize": True, + }, + }, + ) + + with pytest.raises(ValueError, match="not supported with nccl_reshard"): + validate_vllm_quantization_config(generation_config) + + @pytest.mark.parametrize( "field", ["refit_cache_loader_routes"], diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 192b1e96b7f..d74428a14e5 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -224,12 +224,35 @@ def test_process_mxfp8_moe_refit_uses_batched_flashinfer_shuffle( is_mx=True, ) - w13_weight = torch.nn.Parameter(torch.zeros(2, 4, 3), requires_grad=False) - w2_weight = torch.nn.Parameter(torch.zeros(2, 3, 2), requires_grad=False) - w13_scale = torch.nn.Parameter(torch.zeros(2, 4, 1), requires_grad=False) - w2_scale = torch.nn.Parameter(torch.zeros(2, 3, 1), requires_grad=False) - w13_scale_from_checkpoint = torch.ones_like(w13_scale) - w2_scale_from_checkpoint = torch.ones_like(w2_scale) + hidden_size = 512 + intermediate_size = 128 + w13_rows = intermediate_size * (2 if is_gated else 1) + w13_weight = torch.nn.Parameter( + torch.zeros(2, w13_rows, hidden_size, dtype=torch.float8_e4m3fn), + requires_grad=False, + ) + w2_weight = torch.nn.Parameter( + torch.zeros(2, hidden_size, intermediate_size, dtype=torch.float8_e4m3fn), + requires_grad=False, + ) + w13_scale = torch.nn.Parameter( + torch.zeros(2, w13_rows * (hidden_size // 32), dtype=torch.uint8), + requires_grad=False, + ) + w2_scale = torch.nn.Parameter( + torch.zeros(2, hidden_size * (intermediate_size // 32), dtype=torch.uint8), + requires_grad=False, + ) + w13_scale_from_checkpoint = torch.ones( + 2, w13_rows, hidden_size // 32, dtype=torch.uint8 + ) + w2_scale_from_checkpoint = torch.ones( + 2, hidden_size, intermediate_size // 32, dtype=torch.uint8 + ) + moe_config = types.SimpleNamespace( + is_act_and_mul=is_gated, + intermediate_size_per_partition=intermediate_size, + ) layer = types.SimpleNamespace( w13_weight=w13_weight, w2_weight=w2_weight, @@ -241,14 +264,20 @@ def test_process_mxfp8_moe_refit_uses_batched_flashinfer_shuffle( w2_weight_scale_from_checkpoint=types.SimpleNamespace( data=w2_scale_from_checkpoint ), + moe_config=moe_config, ) moe_kernel = object() - moe_quant_config = object() + moe_quant_config = types.SimpleNamespace( + w1_scale=w13_scale, + w2_scale=w2_scale, + ) quant_method = types.SimpleNamespace( - moe=types.SimpleNamespace(is_act_and_mul=is_gated), + moe=moe_config, moe_kernel=moe_kernel, moe_quant_config=moe_quant_config, mxfp8_backend=Fp8MoeBackend.FLASHINFER_TRTLLM, + experts_cls=types.SimpleNamespace(is_monolithic=lambda: True), + weight_block_size=[32, 32], ) shuffled = ( torch.full_like(w13_weight, 1), @@ -564,9 +593,13 @@ def test_load_weights_preserves_prequantized_mxfp8_and_clamps_scales( from nemo_rl.models.generation.vllm import vllm_backend fp8 = fp8_module - fp8.global_fp8_config = types.SimpleNamespace(is_mx=True) + fp8.global_fp8_config = types.SimpleNamespace( + is_mx=True, + refit_prequantize=True, + ) native = torch.ones(2, 2, dtype=torch.bfloat16) prequantized = torch.ones(2, 2, dtype=torch.float8_e4m3fn) + prequantized_scales = torch.ones(2, 1, dtype=torch.uint8) receiver_quantized = torch.full((2, 64), 2.0, dtype=torch.bfloat16) receiver_fp8 = torch.ones(2, 64, dtype=torch.float8_e4m3fn) receiver_scales = torch.tensor([[0, 7], [3, 0]], dtype=torch.uint8) @@ -575,7 +608,7 @@ def test_load_weights_preserves_prequantized_mxfp8_and_clamps_scales( monkeypatch.setattr( fp8, "_is_fp8_weight", - lambda name, _model: name != "model.native", + lambda name, _model: name.endswith(".weight"), ) monkeypatch.setattr( mxfp8_utils, @@ -600,6 +633,10 @@ def test_load_weights_preserves_prequantized_mxfp8_and_clamps_scales( [ ("model.native", native), ("model.prequantized.weight", prequantized), + ( + "model.prequantized.weight_scale_from_checkpoint", + prequantized_scales, + ), ("model.receiver.weight", receiver_quantized), ], types.SimpleNamespace( @@ -612,15 +649,89 @@ def test_load_weights_preserves_prequantized_mxfp8_and_clamps_scales( assert loaded[0][1] is native assert loaded[1][0] == "model.prequantized.weight" assert loaded[1][1] is prequantized - assert loaded[2][0] == "model.receiver.weight" - assert loaded[2][1] is receiver_fp8 - assert loaded[3][0] == "model.receiver.weight_scale_from_checkpoint" + assert loaded[2][0] == "model.prequantized.weight_scale_from_checkpoint" + assert loaded[2][1] is prequantized_scales + assert loaded[3][0] == "model.receiver.weight" + assert loaded[3][1] is receiver_fp8 + assert loaded[4][0] == "model.receiver.weight_scale_from_checkpoint" torch.testing.assert_close( - loaded[3][1], + loaded[4][1], torch.tensor([[1, 7], [3, 1]], dtype=torch.uint8), ) +def test_load_weights_rejects_unnegotiated_mxfp8_payload(fp8_module, monkeypatch): + fp8 = fp8_module + fp8.global_fp8_config = types.SimpleNamespace( + is_mx=True, + refit_prequantize=False, + ) + monkeypatch.setattr(fp8, "_is_fp8_weight", lambda _name, _model: True) + + with pytest.raises(ValueError, match="refit_prequantize=true"): + fp8.load_weights( + [("model.weight", torch.ones(2, 2, dtype=torch.float8_e4m3fn))], + types.SimpleNamespace( + model=object(), + vllm_config=types.SimpleNamespace(additional_config={}), + ), + ) + + +def test_load_weights_rejects_prequantized_mxfp8_without_scale(fp8_module, monkeypatch): + fp8 = fp8_module + fp8.global_fp8_config = types.SimpleNamespace( + is_mx=True, + refit_prequantize=True, + ) + monkeypatch.setattr(fp8, "_is_fp8_weight", lambda _name, _model: True) + + with pytest.raises(ValueError, match="missing.*scale_from_checkpoint"): + fp8.load_weights( + [("model.weight", torch.ones(2, 2, dtype=torch.float8_e4m3fn))], + types.SimpleNamespace( + model=object(), + vllm_config=types.SimpleNamespace(additional_config={}), + ), + ) + + +def test_load_weights_preserves_non_mx_blockwise_fp8_payload(fp8_module, monkeypatch): + from nemo_rl.models.generation.vllm import vllm_backend + + fp8 = fp8_module + fp8.global_fp8_config = types.SimpleNamespace( + is_mx=False, + refit_prequantize=False, + ) + weight = torch.ones(2, 2, dtype=torch.float8_e4m3fn) + scale = torch.ones(1, dtype=torch.float32) + loaded = [] + monkeypatch.setattr( + fp8, + "_is_fp8_weight", + lambda name, _model: name == "model.weight", + ) + monkeypatch.setattr( + vllm_backend, + "load_weights_maybe_cached", + lambda model, weights, *, cache_loader_routes: loaded.extend(weights), + ) + + fp8.load_weights( + [("model.weight", weight), ("model.weight_scale_inv", scale)], + types.SimpleNamespace( + model=object(), + vllm_config=types.SimpleNamespace(additional_config={}), + ), + ) + + assert loaded[0][0] == "model.weight" + assert loaded[0][1] is weight + assert loaded[1][0] == "model.weight_scale_inv" + assert loaded[1][1] is scale + + def test_mxfp8_padding_helpers_preserve_values_and_fill_padding( fp8_module: types.ModuleType, ) -> None: diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index eea32e314bc..258ce45e926 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -515,18 +515,15 @@ def forward(): worker._clear_rope_and_moe_dispatcher_caches() -@pytest.mark.parametrize( - ("selected", "dtype"), - [(False, torch.bfloat16), (True, torch.float8_e4m3fn)], -) -def test_maybe_prequantize_param_passthrough(selected, dtype): +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float8_e4m3fn]) +def test_maybe_prequantize_param_passthrough_when_not_selected(dtype): from nemo_rl.models.policy.workers.megatron_policy_worker import ( MegatronPolicyWorkerImpl, ) worker = object.__new__(MegatronPolicyWorkerImpl) name = "model.weight" - worker._refit_prequant_names = {name} if selected else set() + worker._refit_prequant_names = set() tensor = torch.ones(2, 2, dtype=dtype) result = list(worker._maybe_prequantize_param(name, tensor)) @@ -536,6 +533,35 @@ def test_maybe_prequantize_param_passthrough(selected, dtype): assert result[0][1] is tensor +def test_maybe_prequantize_param_rejects_fp8_trainer_storage(): + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + + worker = object.__new__(MegatronPolicyWorkerImpl) + name = "model.weight" + worker._refit_prequant_names = {name} + tensor = torch.ones(2, 2, dtype=torch.float8_e4m3fn) + + with pytest.raises(ValueError, match="BF16 trainer-exported weights"): + list(worker._maybe_prequantize_param(name, tensor)) + + +def test_enable_refit_prequantize_rejects_blockwise_fp8_storage(): + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + + worker = object.__new__(MegatronPolicyWorkerImpl) + worker.fp8_cfg = { + "fp8_param": True, + "fp8_recipe": "blockwise", + } + + with pytest.raises(ValueError, match="BF16 trainer-exported weights"): + worker.enable_refit_prequantize(["model.weight"]) + + @pytest.mark.parametrize("slim", [False, True]) def test_offload_after_refit_routes_cleanup_by_mode(monkeypatch, slim): from nemo_rl.models.policy.workers.megatron_policy_worker import ( From 94a312e94bb4e03be8cd56f9a003c7d67967aa4f Mon Sep 17 00:00:00 2001 From: seonjinn Date: Mon, 24 Aug 2026 18:38:49 -0700 Subject: [PATCH 33/76] fix: accept serialized_fp8_config in real-quant prepare_refit_info Signed-off-by: seonjinn --- nemo_rl/modelopt/models/generation/vllm_quant_backend.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/nemo_rl/modelopt/models/generation/vllm_quant_backend.py b/nemo_rl/modelopt/models/generation/vllm_quant_backend.py index e3acc58e734..7536a7a89d9 100644 --- a/nemo_rl/modelopt/models/generation/vllm_quant_backend.py +++ b/nemo_rl/modelopt/models/generation/vllm_quant_backend.py @@ -529,10 +529,12 @@ def _synchronize_before_ipc_data_ack(self) -> None: super()._synchronize_before_ipc_data_ack() def prepare_refit_info( - self, state_dict_info: dict[str, Any] + self, + state_dict_info: dict[str, Any], + serialized_fp8_config: Optional[dict[str, Any]] = None, ) -> Optional[list[str]]: if not self._is_real_quant_model(): - return super().prepare_refit_info(state_dict_info) + return super().prepare_refit_info(state_dict_info, serialized_fp8_config) # Real quantization owns a separate refit handshake and must not import # the legacy FP8 quantization path. From 759b44c83b7da7ee625404436356565779520b35 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Mon, 24 Aug 2026 18:38:49 -0700 Subject: [PATCH 34/76] fix: gate slim-refit optimizer offload on offload_optimizer_for_refit Signed-off-by: seonjinn --- .../policy/workers/megatron_policy_worker.py | 1 + tests/unit/models/policy/test_megatron_worker.py | 15 ++++++++++++--- 2 files changed, 13 insertions(+), 3 deletions(-) diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 4ea41466c10..b57766e87bf 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -3055,6 +3055,7 @@ def offload_after_refit(self): hasattr(self, "optimizer") and self.optimizer is not None and not self.optimizer_cpu_offload + and self.offload_optimizer_for_refit ): self.move_optimizer("cpu") gc.collect() diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index ec48b75b856..8a7347fd3d2 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -706,8 +706,13 @@ def test_enable_refit_prequantize_rejects_blockwise_fp8_storage(): worker.enable_refit_prequantize(["model.weight"]) -@pytest.mark.parametrize("slim", [False, True]) -def test_offload_after_refit_routes_cleanup_by_mode(monkeypatch, slim): +@pytest.mark.parametrize( + "slim,offload_optimizer", + [(False, True), (True, True), (True, False)], +) +def test_offload_after_refit_routes_cleanup_by_mode( + monkeypatch, slim, offload_optimizer +): from nemo_rl.models.policy.workers.megatron_policy_worker import ( MegatronPolicyWorkerImpl, ) @@ -727,6 +732,7 @@ def test_offload_after_refit_routes_cleanup_by_mode(monkeypatch, slim): worker._clear_rope_and_moe_dispatcher_caches = MagicMock() worker.optimizer = object() worker.optimizer_cpu_offload = False + worker.offload_optimizer_for_refit = offload_optimizer worker.move_optimizer = MagicMock() worker.offload_before_refit = MagicMock() collect = MagicMock() @@ -753,7 +759,10 @@ def test_offload_after_refit_routes_cleanup_by_mode(monkeypatch, slim): if slim: worker._clear_fp8_caches.assert_called_once_with() worker._clear_rope_and_moe_dispatcher_caches.assert_called_once_with() - worker.move_optimizer.assert_called_once_with("cpu") + if offload_optimizer: + worker.move_optimizer.assert_called_once_with("cpu") + else: + worker.move_optimizer.assert_not_called() collect.assert_called_once_with() empty_cache.assert_called_once_with() worker.offload_before_refit.assert_not_called() From 3736d161ca2e05f6c37911945764010e8511e36d Mon Sep 17 00:00:00 2001 From: seonjinn Date: Mon, 24 Aug 2026 18:38:49 -0700 Subject: [PATCH 35/76] fix: validate MXFP8 scale presence against full refit manifest Signed-off-by: seonjinn --- .../generation/vllm/quantization/fp8.py | 19 ++++++++++++++++++- .../models/generation/vllm/vllm_backend.py | 1 + 2 files changed, 19 insertions(+), 1 deletion(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index 6b4ed803b23..9eebb416c81 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -73,6 +73,11 @@ class FP8State: seen_params: set = field(default_factory=lambda: set()) fp8_param_names: set = field(default_factory=lambda: set()) vllm_patches: list = field(default_factory=lambda: []) + # Full refit manifest names from prepare_refit_info. load_weights receives + # transfer batches split by buffer size, so a weight and its + # *_scale_from_checkpoint entry may arrive in different batches; presence + # must be validated against the full manifest, not the current batch. + refit_manifest_names: set | None = None # Global FP8 config that can be accessed by patched vLLM functions @@ -112,6 +117,10 @@ def install_fp8_config(config: dict[str, Any] | None) -> None: global_fp8_config = FP8Config(**config) +def set_refit_manifest_names(names: set[str] | None) -> None: + fp8_state.refit_manifest_names = names + + def _patch_ray_executor_v2_worker(ray_executor_v2: Any, fp8_config: FP8Config) -> None: """Install FP8 patches inside RayExecutorV2 workers before model loading.""" original_ray_worker_proc = ray_executor_v2.RayWorkerProc @@ -529,8 +538,16 @@ def load_weights(weights, model_runner): "MXFP8 E4M3 refit weights require refit_prequantize=true; " "other FP8 trainer scale layouts are not compatible." ) + # Transfer batches split by buffer size, so the matching scale may + # arrive in an earlier or later batch; validate against the full + # refit manifest when available, not just the current batch. scale_name = k + "_scale_from_checkpoint" - if global_fp8_config.is_mx and scale_name not in weight_names: + manifest = fp8_state.refit_manifest_names + if ( + global_fp8_config.is_mx + and scale_name not in weight_names + and (manifest is None or scale_name not in manifest) + ): raise ValueError( f"Prequantized MXFP8 weight {k!r} is missing {scale_name!r}." ) diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index f64dc451e9f..c30984fe9a8 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -480,6 +480,7 @@ def prepare_refit_info( from nemo_rl.models.generation.vllm.quantization import fp8 fp8.install_fp8_config(serialized_fp8_config) + fp8.set_refit_manifest_names(set(state_dict_info)) if not ( fp8.global_fp8_config is not None and fp8.global_fp8_config.is_mx From 8dc4c367e4306ba69bf4e244e6c92d2cd91cbac5 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Tue, 25 Aug 2026 23:59:03 -0700 Subject: [PATCH 36/76] fix(dynamo): align prepare_refit_info override with interface Signed-off-by: seonjinn --- nemo_rl/models/generation/dynamo/dynamo_generation.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/nemo_rl/models/generation/dynamo/dynamo_generation.py b/nemo_rl/models/generation/dynamo/dynamo_generation.py index 1c488855c7d..4f338e7b048 100644 --- a/nemo_rl/models/generation/dynamo/dynamo_generation.py +++ b/nemo_rl/models/generation/dynamo/dynamo_generation.py @@ -621,12 +621,17 @@ def init_collective( train_world_size=train_world_size, ) - def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: + def prepare_refit_info( + self, state_dict_info: Optional[dict[str, Any]] + ) -> Optional[list[str]]: """Serialize checkpoint-format tensor metadata for native vLLM refit.""" + if state_dict_info is None: + return None channel = self._refit_channel if channel is None: raise RuntimeError("Dynamo refit channel is unavailable") channel.prepare(state_dict_info) + return None def update_weights_via_ipc_zmq(self) -> list[ray.ObjectRef]: raise NotImplementedError( From c162cfd11d73ff4d8753347131f483c9416f31ca Mon Sep 17 00:00:00 2001 From: seonjinn Date: Wed, 26 Aug 2026 12:04:18 -0700 Subject: [PATCH 37/76] fix: resolve CI unit-test failures after main merge - Guard offload_after_refit against configs without megatron_cfg - Skip fp8 module import in prepare_refit_info for non-FP8 refits so stubbed quant-backend tests can prepare refit info - Drop duplicate MXFP8 scale clamp already done in quantize_mxfp8_weight - Fix prequantized-load test to compare tensor contents (reshape breaks object identity) and pin the noncolocated PPO mock's refit negotiation Signed-off-by: seonjinn --- nemo_rl/models/generation/vllm/quantization/fp8.py | 6 ------ nemo_rl/models/generation/vllm/vllm_backend.py | 6 ++++++ nemo_rl/models/policy/workers/megatron_policy_worker.py | 2 +- tests/unit/algorithms/test_ppo.py | 1 + tests/unit/models/generation/test_vllm_fp8_quantization.py | 5 ++++- 5 files changed, 12 insertions(+), 8 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index 0dab18e0c6e..b8e5d8124c2 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -668,12 +668,6 @@ def load_weights(weights, model_runner): weight_block_size=FP8_BLOCK_QUANT_KWARGS["weight_block_size"], ) if global_fp8_config.is_mx: - # vLLM 0.25 returns row-major [M, K / 32] E8M0 scales. - # All-zero blocks quantize to E8M0 byte 0, which destabilizes the - # TRTLLM MXFP8 kernel; clamp to byte 1 (weights are 0 anyway). - param_scale = torch.where( - param_scale == 0, torch.ones_like(param_scale), param_scale - ) weights_quantized.append([k, param_lp]) weights_quantized.append([k + "_scale_from_checkpoint", param_scale]) else: diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index e649a2879e0..dd4574127d2 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -580,6 +580,12 @@ def prepare_refit_info( self._validate_native_layerwise_refit() self.state_dict_info = state_dict_info # pyrefly: ignore[implicitly-defined-attribute] This class does not define __init__ so assignments like this should be ignored + # Non-FP8 runs serialize no FP8 config; skip the fp8 module (and its + # heavyweight vLLM imports) entirely so quant backends stubbed without + # the full vLLM surface can still prepare refit info. + if serialized_fp8_config is None: + return None + from nemo_rl.models.generation.vllm.quantization import fp8 fp8.install_fp8_config(serialized_fp8_config) diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 98b08a6359c..22e35de22d5 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -3495,7 +3495,7 @@ def offload_after_refit(self): ) self.model.eval() torch.randn(1).cuda() # wake up torch allocator - if self.cfg["megatron_cfg"].get("refit_slim_offload_after"): + if self.cfg.get("megatron_cfg", {}).get("refit_slim_offload_after"): # Grad buffers were already offloaded by offload_before_refit at # the start of the refit, so skip the full rerun (grad-buffer moves # and a second gc/empty_cache pair). Cache clears must still honor diff --git a/tests/unit/algorithms/test_ppo.py b/tests/unit/algorithms/test_ppo.py index 8f51fa4f9fa..96965236c29 100644 --- a/tests/unit/algorithms/test_ppo.py +++ b/tests/unit/algorithms/test_ppo.py @@ -1462,6 +1462,7 @@ def get_placement_groups(self): policy.init_collective.return_value = ["policy-future"] value_model = MagicMock() generation = MagicMock() + generation.prepare_refit_info.return_value = None generation.init_collective.return_value = ["generation-future"] policy_factory = MagicMock(return_value=policy) value_factory = MagicMock(return_value=value_model) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index cbe4b02907e..0e10f28eb96 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -1094,7 +1094,10 @@ def test_load_weights_preserves_prequantized_mxfp8_and_clamps_scales( assert loaded[2][0] == "model.prequantized.weight_scale_from_checkpoint" assert loaded[2][1] is prequantized_scales assert loaded[3][0] == "model.receiver.weight" - assert loaded[3][1] is receiver_fp8 + # quantize_mxfp8_weight reshapes to the checkpoint layout, so compare + # contents rather than object identity. + assert loaded[3][1].dtype == torch.float8_e4m3fn + assert torch.equal(loaded[3][1].view(torch.uint8), receiver_fp8.view(torch.uint8)) assert loaded[4][0] == "model.receiver.weight_scale_from_checkpoint" torch.testing.assert_close( loaded[4][1], From e28239c57b3174ca359362c1cea7bc710e19e631 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Wed, 26 Aug 2026 15:11:07 -0700 Subject: [PATCH 38/76] fix(refit): derive prequantize metadata without re-exporting weights Signed-off-by: seonjinn --- .../vllm/quantization/fp8_train_utils.py | 2 +- .../policy/workers/megatron_policy_worker.py | 45 ++++++++++++- .../models/policy/test_megatron_worker.py | 66 ++++++++++++++++++- 3 files changed, 108 insertions(+), 5 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py index f7aa25f91ce..a8e55fc3b70 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py @@ -90,7 +90,7 @@ def mxfp8_e4m3_quantize_for_refit( # Match the receiver path's zero-scale clamp: an E8M0 byte of 0 (2^-127) # destabilizes the TRTLLM kernels, and pre-quantized tensors skip the # receiver-side quantize branch where the clamp normally runs. - x_scales = torch.where(x_scales == 0, torch.ones_like(x_scales), x_scales) + x_scales = x_scales.masked_fill(x_scales == 0, 1) return x_q, x_scales diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 22e35de22d5..3bfe4929914 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -468,6 +468,12 @@ def __init__( # HF param names to MXFP8-quantize on the trainer during refit; set via # enable_refit_prequantize() when vllm_cfg.refit_prequantize is on. self._refit_prequant_names: set[str] = set() + # HF param metadata cached by prepare_refit_info so that + # enable_refit_prequantize can derive updated metadata without + # re-running the export + quantize pass. + self._refit_param_info_hf: Optional[ + dict[str, tuple[torch.Size, torch.dtype]] + ] = None # Pinned host staging for the reference-policy swap; only populated when # megatron_cfg["pinned_reference_swap"] is enabled. Buffer contents are # only live within a single use_reference_model call (every copy @@ -2211,7 +2217,7 @@ def get_topk_logits( @torch.no_grad() @wrap_with_nvtx_name("megatron_policy_worker/prepare_refit_info") - def prepare_refit_info(self) -> None: + def prepare_refit_info(self) -> Optional[dict[str, Any]]: """Prepare state dict metadata for weight refitting and IPC streaming.""" self.refit_param_info_mcore = self._calculate_refit_param_info() @@ -2220,11 +2226,17 @@ def prepare_refit_info(self) -> None: for name, tensor in self._iter_params_with_optional_kv_scales(): refit_param_info_hf[name] = (tensor.shape, tensor.dtype) + self._refit_param_info_hf = refit_param_info_hf return refit_param_info_hf def enable_refit_prequantize(self, param_names: list[str]) -> dict[str, Any]: """Quantize the listed HF params to MXFP8 on the trainer during refit. + Derives the updated metadata from the shapes cached by + prepare_refit_info instead of re-running the export + quantize pass: + MXFP8 keeps the value shape and adds one uint8 E8M0 scale per 32-wide + block along the last dim, so no tensor needs to be materialized here. + Args: param_names: fp8-eligible parameter names reported by the vLLM workers (see VllmInternalWorkerExtension.prepare_refit_info). @@ -2238,12 +2250,39 @@ def enable_refit_prequantize(self, param_names: list[str]) -> dict[str, Any]: "vllm_cfg.refit_prequantize requires BF16 trainer-exported weights; " "Megatron blockwise FP8 parameter storage uses a different scale layout." ) + if self._refit_param_info_hf is None: + raise RuntimeError( + "enable_refit_prequantize requires prepare_refit_info to have " + "run first so the HF parameter metadata is available." + ) + + from nemo_rl.models.generation.vllm.quantization.fp8_train_utils import ( + MXFP8_BLOCK_SIZE, + ) self._refit_prequant_names = set(param_names) refit_param_info_hf = {} - for name, tensor in self._iter_params_with_optional_kv_scales(): - refit_param_info_hf[name] = (tensor.shape, tensor.dtype) + for name, (shape, dtype) in self._refit_param_info_hf.items(): + if name not in self._refit_prequant_names: + refit_param_info_hf[name] = (shape, dtype) + continue + if dtype == torch.float8_e4m3fn: + raise ValueError( + "vllm_cfg.refit_prequantize requires BF16 trainer-exported weights; " + f"{name} is already stored as E4M3 with a non-MXFP8 scale layout." + ) + if shape[-1] % MXFP8_BLOCK_SIZE != 0: + raise ValueError( + f"MXFP8 requires the last dim to be divisible by " + f"{MXFP8_BLOCK_SIZE}; {name} has shape {tuple(shape)}." + ) + scale_shape = torch.Size((*shape[:-1], shape[-1] // MXFP8_BLOCK_SIZE)) + refit_param_info_hf[name] = (shape, torch.float8_e4m3fn) + refit_param_info_hf[name + "_scale_from_checkpoint"] = ( + scale_shape, + torch.uint8, + ) return refit_param_info_hf def _maybe_prequantize_param( diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index a08d9b22a8e..ee39f110fd1 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -508,6 +508,7 @@ def test_checkpoint_engine_prequant_handshake_exports_mxfp8_weights(): class _PrequantCheckpointWorker(MegatronCheckpointEngineSendMixin): enable_refit_prequantize = MegatronPolicyWorkerImpl.enable_refit_prequantize + _is_fp8_export = MegatronPolicyWorkerImpl._is_fp8_export _iter_params_with_optional_kv_scales = ( MegatronPolicyWorkerImpl._iter_params_with_optional_kv_scales ) @@ -517,6 +518,8 @@ class _PrequantCheckpointWorker(MegatronCheckpointEngineSendMixin): weight = torch.randn(64, 64, dtype=torch.bfloat16) worker = _PrequantCheckpointWorker() worker._refit_prequant_names = set() + worker._refit_param_info_hf = None + worker.fp8_cfg = None worker.model = object() worker.draft_model = None worker.refit_conversion_tasks = [] @@ -524,7 +527,12 @@ class _PrequantCheckpointWorker(MegatronCheckpointEngineSendMixin): worker.megatron_bridge = SimpleNamespace( export_hf_weights=lambda *_args, **_kwargs: iter([(name, weight)]) ) - worker.prepare_refit_info = lambda: {name: (weight.shape, weight.dtype)} + + def _prepare_refit_info() -> dict[str, Any]: + worker._refit_param_info_hf = {name: (weight.shape, weight.dtype)} + return worker._refit_param_info_hf + + worker.prepare_refit_info = _prepare_refit_info worker.checkpoint_engine = SimpleNamespace(get_target_weight_layout=lambda: None) class _Generation: @@ -733,6 +741,62 @@ def test_enable_refit_prequantize_rejects_blockwise_fp8_storage(): worker.enable_refit_prequantize(["model.weight"]) +def test_enable_refit_prequantize_requires_prepare_refit_info(): + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + + worker = object.__new__(MegatronPolicyWorkerImpl) + worker.fp8_cfg = None + worker._refit_param_info_hf = None + + with pytest.raises(RuntimeError, match="prepare_refit_info"): + worker.enable_refit_prequantize(["model.weight"]) + + +def test_enable_refit_prequantize_derives_metadata_without_export(): + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + + worker = object.__new__(MegatronPolicyWorkerImpl) + worker.fp8_cfg = None + worker._refit_param_info_hf = { + "model.a.weight": (torch.Size([4, 64]), torch.bfloat16), + "model.b.weight": (torch.Size([4, 64]), torch.bfloat16), + } + + def _fail_iter(*_args, **_kwargs): + raise AssertionError("metadata derivation must not re-export weights") + + worker._iter_params_with_optional_kv_scales = _fail_iter + + info = worker.enable_refit_prequantize(["model.a.weight"]) + + assert info["model.a.weight"] == (torch.Size([4, 64]), torch.float8_e4m3fn) + assert info["model.a.weight_scale_from_checkpoint"] == ( + torch.Size([4, 2]), + torch.uint8, + ) + assert info["model.b.weight"] == (torch.Size([4, 64]), torch.bfloat16) + assert worker._refit_prequant_names == {"model.a.weight"} + + +def test_enable_refit_prequantize_rejects_indivisible_last_dim(): + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + + worker = object.__new__(MegatronPolicyWorkerImpl) + worker.fp8_cfg = None + worker._refit_param_info_hf = { + "model.weight": (torch.Size([4, 48]), torch.bfloat16), + } + + with pytest.raises(ValueError, match="divisible"): + worker.enable_refit_prequantize(["model.weight"]) + + @pytest.mark.parametrize( "slim,offload_optimizer", [(False, True), (True, True), (True, False)], From 232da176d1758ae167dfc09a28ae0fa14bcfbe92 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Wed, 26 Aug 2026 16:03:53 -0700 Subject: [PATCH 39/76] Fix pyrefly no-matching-overload false positive on masked_fill Signed-off-by: seonjinn --- nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py | 1 + 1 file changed, 1 insertion(+) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py index a8e55fc3b70..c4d2df7e290 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py @@ -90,6 +90,7 @@ def mxfp8_e4m3_quantize_for_refit( # Match the receiver path's zero-scale clamp: an E8M0 byte of 0 (2^-127) # destabilizes the TRTLLM kernels, and pre-quantized tensors skip the # receiver-side quantize branch where the clamp normally runs. + # pyrefly: ignore # no-matching-overload x_scales = x_scales.masked_fill(x_scales == 0, 1) return x_q, x_scales From 6b2ab6ef746f61025ee52111b806d93980c6632d Mon Sep 17 00:00:00 2001 From: seonjinn Date: Wed, 26 Aug 2026 18:05:55 -0700 Subject: [PATCH 40/76] test: update refit cleanup fixture after main merge Signed-off-by: seonjinn --- tests/unit/models/policy/test_megatron_worker.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index ee39f110fd1..c868f52a2b2 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -826,6 +826,7 @@ def test_offload_after_refit_routes_cleanup_by_mode( worker.offload_optimizer_for_refit = offload_optimizer worker.move_optimizer = MagicMock() worker.offload_before_refit = MagicMock() + worker.finalize_async_save = MagicMock() collect = MagicMock() empty_cache = MagicMock() monkeypatch.setattr( @@ -845,6 +846,7 @@ def test_offload_after_refit_routes_cleanup_by_mode( worker.offload_after_refit() + worker.finalize_async_save.assert_called_once_with() worker.move_model.assert_called_once_with(model, "cpu") model.eval.assert_called_once_with() if slim: From 06b6fbb08a6ac95d1942eb0f7bedbd2fb659fefe Mon Sep 17 00:00:00 2001 From: seonjinn Date: Wed, 26 Aug 2026 23:22:30 -0700 Subject: [PATCH 41/76] refactor(refit): limit optimized recipes to sync RL Signed-off-by: seonjinn --- ...3-235b-16n8g-async-1off-mxfp8-rollout.yaml | 17 ----- ...3-235b-32n4g-async-1off-mxfp8-rollout.yaml | 2 - ...-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml | 4 +- ...-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml | 20 ------ ...en3-32b-8n4g-async-1off-mxfp8-rollout.yaml | 2 - tests/test_mxfp8_rollout_recipes.py | 68 ++++++++----------- ...en3-235b-16n8g-async-1off-mxfp8-rollout.sh | 48 ------------- ...n3-30ba3b-4n8g-async-1off-mxfp8-rollout.sh | 43 ------------ tests/test_suites/performance.txt | 5 -- 9 files changed, 31 insertions(+), 178 deletions(-) delete mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml delete mode 100644 examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml delete mode 100755 tests/test_suites/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.sh delete mode 100755 tests/test_suites/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.sh diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml deleted file mode 100644 index 2296cfd580f..00000000000 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.yaml +++ /dev/null @@ -1,17 +0,0 @@ -defaults: ./grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml -checkpointing: - checkpoint_dir: results/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout # pragma: allowlist secret -policy: - generation: - colocated: - resources: - num_nodes: 8 - gpus_per_node: 8 -logger: - log_dir: logs/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout - wandb: - name: grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout -cluster: - gpus_per_node: 8 - num_nodes: 16 - segment_size: 8 diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml index 5ffdafc75c3..06e59d18399 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout.yaml @@ -7,8 +7,6 @@ policy: tensor_parallel_size: 4 precision: "fp8" is_mx: true - refit_prequantize: true - refit_cache_loader_routes: true quantization_ignore_patterns: - model.layers.*.self_attn.* - model.layers.*.mlp.gate diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml index 09824874f5b..bf8376355d3 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout.yaml @@ -21,8 +21,8 @@ policy: gpu_memory_utilization: 0.8 precision: "fp8" is_mx: true - refit_prequantize: true - refit_cache_loader_routes: true + refit_prequantize: false + refit_cache_loader_routes: false quantization_ignore_patterns: - model.layers.*.self_attn.* - model.layers.*.mlp.gate diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml deleted file mode 100644 index 4a3dcead8b2..00000000000 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.yaml +++ /dev/null @@ -1,20 +0,0 @@ -defaults: ./grpo-qwen3-30ba3b-4n8g-async-1off.yaml -checkpointing: - checkpoint_dir: results/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout -policy: - generation: - vllm_cfg: - precision: "fp8" - is_mx: true - refit_prequantize: true - refit_cache_loader_routes: true - quantization_ignore_patterns: - - model.layers.*.self_attn.* - - model.layers.*.mlp.gate - - lm_head - vllm_kwargs: - moe_backend: flashinfer_trtllm -logger: - log_dir: logs/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout - wandb: - name: grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout diff --git a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-8n4g-async-1off-mxfp8-rollout.yaml b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-8n4g-async-1off-mxfp8-rollout.yaml index 3ad80cff7f4..5f547d6516f 100644 --- a/examples/configs/recipes/llm/performance/grpo-qwen3-32b-8n4g-async-1off-mxfp8-rollout.yaml +++ b/examples/configs/recipes/llm/performance/grpo-qwen3-32b-8n4g-async-1off-mxfp8-rollout.yaml @@ -9,8 +9,6 @@ policy: gpu_memory_utilization: 0.8 precision: "fp8" is_mx: true - refit_prequantize: true - refit_cache_loader_routes: true quantization_ignore_patterns: - model.layers.*.self_attn.* - lm_head diff --git a/tests/test_mxfp8_rollout_recipes.py b/tests/test_mxfp8_rollout_recipes.py index 0a028f780e9..7d1d2d8b14e 100644 --- a/tests/test_mxfp8_rollout_recipes.py +++ b/tests/test_mxfp8_rollout_recipes.py @@ -12,7 +12,6 @@ # See the License for the specific language governing permissions and # limitations under the License. -from fnmatch import fnmatch from pathlib import Path import pytest @@ -129,7 +128,6 @@ "segment_size": 2, "async_engine": True, "moe_backend": "flashinfer_trtllm", - "refit_optimizations": True, "train_global_batch_size": 2048, "ignore_patterns": [ "model.layers.*.self_attn.*", @@ -155,7 +153,6 @@ "segment_size": 4, "async_engine": True, "moe_backend": "flashinfer_trtllm", - "refit_optimizations": True, "ignore_patterns": [ "model.layers.*.self_attn.*", "lm_head", @@ -182,7 +179,6 @@ "async_engine": True, "tensor_parallel_size": 4, "moe_backend": "flashinfer_trtllm", - "refit_optimizations": True, "ignore_patterns": [ "model.layers.*.self_attn.*", "model.layers.*.mlp.gate", @@ -191,13 +187,6 @@ }, } -B200_MXFP8_RECIPES = frozenset( - { - "grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout", - "grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout", - } -) - REMOVED_BACKEND_VARIANT_SUFFIXES = ( "-flashinfer-trtllm", "-triton", @@ -295,22 +284,44 @@ def test_mxfp8_rollout_recipe_matrix(case_name: str, expected: dict) -> None: @pytest.mark.parametrize( - "config_path", - sorted(PERF_CONFIG_DIR.glob("grpo-qwen3-*mxfp8-rollout.yaml")), - ids=lambda path: path.stem, + "case_name", + ( + "grpo-qwen3-30ba3b-4n4g-mxfp8-rollout", + "grpo-qwen3-32b-4n4g-mxfp8-rollout", + "grpo-qwen3-235b-16n4g-mxfp8-rollout", + ), ) -def test_all_qwen3_mxfp8_rollout_recipes_enable_refit_optimizations( - config_path: Path, +def test_sync_qwen3_mxfp8_rollout_recipes_enable_refit_optimizations( + case_name: str, ) -> None: - config = _load_resolved_yaml(config_path) + config = _load_resolved_yaml(PERF_CONFIG_DIR / f"{case_name}.yaml") vllm_cfg = config["policy"]["generation"]["vllm_cfg"] + assert vllm_cfg["refit_prequantize"] is True assert vllm_cfg["refit_cache_loader_routes"] is True +@pytest.mark.parametrize( + "case_name", + ( + "grpo-qwen3-30ba3b-4n4g-async-1off-mxfp8-rollout", + "grpo-qwen3-32b-8n4g-async-1off-mxfp8-rollout", + "grpo-qwen3-235b-32n4g-async-1off-mxfp8-rollout", + ), +) +def test_async_qwen3_mxfp8_rollout_recipes_skip_sync_refit_optimizations( + case_name: str, +) -> None: + config = _load_resolved_yaml(PERF_CONFIG_DIR / f"{case_name}.yaml") + vllm_cfg = config["policy"]["generation"]["vllm_cfg"] + + assert not vllm_cfg.get("refit_prequantize", False) + assert not vllm_cfg.get("refit_cache_loader_routes", False) + + def test_mxfp8_rollout_recipes_are_in_gb200_performance_suite() -> None: recipe_names = {path.stem for path in PERF_CONFIG_DIR.glob("*-mxfp8-rollout.yaml")} - assert recipe_names == set(MXFP8_CASES) | set(B200_MXFP8_RECIPES) + assert recipe_names == set(MXFP8_CASES) suite_text = GB200_SUITE.read_text(encoding="utf-8") @@ -365,27 +376,6 @@ def test_qwen3_235b_mxfp8_recipes_keep_baseline_runtime_knobs() -> None: assert "max_val_samples" not in grpo_config -@pytest.mark.parametrize("case_name", sorted(B200_MXFP8_RECIPES)) -def test_b200_async_mxfp8_recipes_keep_router_gate_in_bf16(case_name: str) -> None: - config = _load_resolved_yaml(PERF_CONFIG_DIR / f"{case_name}.yaml") - vllm_cfg = config["policy"]["generation"]["vllm_cfg"] - patterns = vllm_cfg["quantization_ignore_patterns"] - - assert vllm_cfg["refit_prequantize"] is True - assert vllm_cfg["refit_cache_loader_routes"] is True - assert patterns == [ - "model.layers.*.self_attn.*", - "model.layers.*.mlp.gate", - "lm_head", - ] - router_name = "model.layers.0.mlp.gate" - expert_name = "model.layers.0.mlp.experts.0.gate_proj" - fused_expert_name = "model.layers.0.mlp.experts.gate_up_proj" - assert any(fnmatch(router_name, pattern) for pattern in patterns) - assert not any(fnmatch(expert_name, pattern) for pattern in patterns) - assert not any(fnmatch(fused_expert_name, pattern) for pattern in patterns) - - def test_deepseek_mxfp8_launchers_handle_unset_and_spaced_checkpoint_paths() -> None: for case_name in ( "grpo-deepseek-v3-64n4g-mxfp8-rollout", diff --git a/tests/test_suites/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.sh b/tests/test_suites/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.sh deleted file mode 100755 index 54879672503..00000000000 --- a/tests/test_suites/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.sh +++ /dev/null @@ -1,48 +0,0 @@ -#!/bin/bash -SCRIPT_DIR=$( cd -- "$( dirname -- "${BASH_SOURCE[0]}" )" &> /dev/null && pwd) -source $SCRIPT_DIR/common.env -# disable NVLS to avoid OOM issue -export NCCL_NVLS_ENABLE=0 -export RAY_CGRAPH_get_timeout=2400 - -# ===== BEGIN CONFIG ===== -NUM_NODES=16 -GPUS_PER_NODE=8 -SEGMENT_SIZE=8 -STEPS_PER_RUN=10 -MAX_STEPS=10 -NUM_RUNS=$(( (MAX_STEPS + STEPS_PER_RUN - 1) / STEPS_PER_RUN )) # Round up -NUM_MINUTES=100 -# ===== END CONFIG ===== - -exit_if_max_steps_reached - -# Run the experiment -cd $PROJECT_ROOT -uv run examples/run_grpo.py \ - --config $CONFIG_PATH \ - grpo.max_num_steps=$MAX_STEPS \ - logger.log_dir=$LOG_DIR \ - logger.wandb_enabled=True \ - logger.wandb.project=nemo-rl \ - logger.wandb.name=$EXP_NAME \ - logger.monitor_gpus=True \ - logger.tensorboard_enabled=True \ - checkpointing.enabled=True \ - checkpointing.checkpoint_dir=$CKPT_DIR \ - +policy.generation.vllm_kwargs.distributed_timeout_seconds=2400 \ - $@ \ - 2>&1 | tee $RUN_LOG - -# Convert tensorboard logs to json -uv run tests/json_dump_tb_logs.py $LOG_DIR --output_path $JSON_METRICS - -# Only run metrics if the target step is reached -if [[ $(jq 'to_entries | .[] | select(.key == "train/loss") | .value | keys | map(tonumber) | max' $JSON_METRICS) -ge $MAX_STEPS ]]; then - uv run tests/check_metrics.py $JSON_METRICS \ - 'median(data["train/token_mult_prob_error"]) < 1.1' \ - 'data["train/token_mult_prob_error"]["10"] < 1.1' - - # Clean up checkpoint directory after successful run to save space. - rm -rf "$CKPT_DIR" -fi diff --git a/tests/test_suites/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.sh b/tests/test_suites/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.sh deleted file mode 100755 index 75e1bc4db31..00000000000 --- a/tests/test_suites/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.sh +++ /dev/null @@ -1,43 +0,0 @@ -#!/bin/bash -SCRIPT_DIR=$( cd -- "$( dirname -- "${BASH_SOURCE[0]}" )" &> /dev/null && pwd) -source $SCRIPT_DIR/common.env - -# ===== BEGIN CONFIG ===== -NUM_NODES=4 -GPUS_PER_NODE=8 -STEPS_PER_RUN=10 -MAX_STEPS=10 -NUM_RUNS=$(( (MAX_STEPS + STEPS_PER_RUN - 1) / STEPS_PER_RUN )) # Round up -NUM_MINUTES=100 -# ===== END CONFIG ===== - -exit_if_max_steps_reached - -# Run the experiment -cd $PROJECT_ROOT -uv run examples/run_grpo.py \ - --config $CONFIG_PATH \ - grpo.max_num_steps=$MAX_STEPS \ - logger.log_dir=$LOG_DIR \ - logger.wandb_enabled=True \ - logger.wandb.project=nemo-rl \ - logger.wandb.name=$EXP_NAME \ - logger.monitor_gpus=True \ - logger.tensorboard_enabled=True \ - checkpointing.enabled=True \ - checkpointing.checkpoint_dir=$CKPT_DIR \ - $@ \ - 2>&1 | tee $RUN_LOG - -# Convert tensorboard logs to json -uv run tests/json_dump_tb_logs.py $LOG_DIR --output_path $JSON_METRICS - -# Only run metrics if the target step is reached -if [[ $(jq 'to_entries | .[] | select(.key == "train/loss") | .value | keys | map(tonumber) | max' $JSON_METRICS) -ge $MAX_STEPS ]]; then - uv run tests/check_metrics.py $JSON_METRICS \ - 'median(data["train/token_mult_prob_error"]) < 1.1' \ - 'data["train/token_mult_prob_error"]["10"] < 1.1' - - # Clean up checkpoint directory after successful run to save space. - rm -rf "$CKPT_DIR" -fi diff --git a/tests/test_suites/performance.txt b/tests/test_suites/performance.txt index 7eaa3dea317..52b3995859b 100644 --- a/tests/test_suites/performance.txt +++ b/tests/test_suites/performance.txt @@ -33,8 +33,3 @@ tests/test_suites/llm/performance/grpo-qwen3-30ba3b-24n8g-async-8off.sh tests/test_suites/llm/performance/grpo-deepseek-v3-64n8g-fp8-async-1off.sh tests/test_suites/llm/performance/grpo-llama3.1-8b-instruct-2n8g-fp8-async-1off.sh -# B200 MXFP8 - -## ASYNC 1-off -tests/test_suites/llm/performance/grpo-qwen3-30ba3b-4n8g-async-1off-mxfp8-rollout.sh -tests/test_suites/llm/performance/grpo-qwen3-235b-16n8g-async-1off-mxfp8-rollout.sh From bf03429727b9c1274b894358450888c0c6a3be86 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 4 Sep 2026 05:01:27 -0700 Subject: [PATCH 42/76] fix(mxfp8): handle mixed grouped expert refit Signed-off-by: seonjinn --- .../generation/vllm/quantization/fp8.py | 54 +++---------- .../generation/test_vllm_fp8_quantization.py | 81 ++++++------------- 2 files changed, 38 insertions(+), 97 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index 293ca553027..029025acea1 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -126,45 +126,7 @@ def set_refit_manifest_names(names: set[str] | None) -> None: fp8_state.refit_manifest_names = names -def _patch_ray_executor_v2_worker(ray_executor_v2: Any, fp8_config: FP8Config) -> None: - """Install FP8 patches inside RayExecutorV2 workers before model loading.""" - original_ray_worker_proc = ray_executor_v2.RayWorkerProc - if getattr(original_ray_worker_proc, "_nrl_fp8_patched", False): - original_ray_worker_proc._nrl_fp8_config = fp8_config - return - - class NRLFP8RayWorkerProc(original_ray_worker_proc): - _nrl_fp8_patched = True - _nrl_fp8_config = fp8_config - - def initialize_worker( - self, - local_rank: int, - env_vars: dict[str, str], - driver_env_vars: dict[str, str] | None = None, - assigned_physical_gpu_ids: list[int] | None = None, - ) -> Any: - global fp8_patches_applied - if not fp8_patches_applied: - apply_fp8_patches(None, type(self)._nrl_fp8_config) - return super().initialize_worker( - local_rank, - env_vars, - driver_env_vars, - assigned_physical_gpu_ids, - ) - - ray_executor_v2.RayWorkerProc = NRLFP8RayWorkerProc - - def monkey_patch_vllm_ray_executor(fp8_config): - try: - from vllm.v1.executor import ray_executor_v2 - except ImportError: - pass - else: - _patch_ray_executor_v2_worker(ray_executor_v2, fp8_config) - if fp8_config.model_parallel_size > 1: if envs.VLLM_USE_RAY_V2_EXECUTOR_BACKEND: from vllm.v1.executor.ray_executor_v2 import RayWorkerProc @@ -573,11 +535,21 @@ def _get_module_from_param_name(model, name: str): return current_module +_GROUPED_EXPERT_WEIGHT_SUFFIXES = ( + "mlp.experts.gate_up_proj", + "mlp.experts.down_proj", +) + + +def _is_grouped_expert_weight(name: str) -> bool: + return name.endswith(_GROUPED_EXPERT_WEIGHT_SUFFIXES) + + def _is_fp8_weight(name, model): if name not in fp8_state.seen_params: fp8_state.seen_params.add(name) # Filter out bias params - if name.endswith("weight"): + if name.endswith("weight") or _is_grouped_expert_weight(name): module = _get_module_from_param_name(model, name) # We currently only quantize linear layers if ( @@ -625,9 +597,7 @@ def load_weights(weights, model_runner): # load their per-block scales. Expand them into the per-expert FP8 (w13, w2 -> w1, w2, and w3) # layout, then reshape to 2D [num_experts, out_features, in_features] -> [num_experts*out_features, in_features] # so the block scales can be quantized and routed correctly. - if k.endswith("mlp.experts.gate_up_proj") or k.endswith( - "mlp.experts.down_proj" - ): + if _is_grouped_expert_weight(k): # Quantize only if vLLM built this layer's experts as FP8. Experts # covered by ``ignored_layers`` (num_{first,last}_layers_in_bf16 / # quantization_ignored_layer_kws) are built unquantized, with bf16 diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 707b2ffa776..f15268ce11c 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -16,7 +16,6 @@ from pathlib import Path from typing import Any -import cloudpickle import pytest import torch import yaml @@ -94,60 +93,6 @@ def test_init_fp8_uses_mxfp8_quantization_config(fp8_module, monkeypatch): assert "VLLM_USE_DEEP_GEMM_E8M0" not in fp8.os.environ -def test_ray_executor_v2_worker_applies_fp8_patches_before_model_load( - fp8_module, monkeypatch -): - fp8 = fp8_module - monkeypatch.setattr(fp8, "_test_applied_configs", [], raising=False) - config = fp8.FP8Config( - use_fp8_weights=True, - model_parallel_size=2, - is_mx=True, - ) - - class FakeRayWorkerProc: - def initialize_worker( - self, - local_rank, - env_vars, - driver_env_vars=None, - assigned_physical_gpu_ids=None, - ): - assert fp8.fp8_patches_applied - return ( - local_rank, - env_vars, - driver_env_vars, - assigned_physical_gpu_ids, - ) - - def fake_apply_fp8_patches(_self, fp8_config): - fp8._test_applied_configs.append(fp8_config) - fp8.fp8_patches_applied = True - - monkeypatch.setattr(fp8, "apply_fp8_patches", fake_apply_fp8_patches) - ray_executor_v2 = types.SimpleNamespace(RayWorkerProc=FakeRayWorkerProc) - fp8._patch_ray_executor_v2_worker(ray_executor_v2, config) - patched_worker_cls = cloudpickle.loads( - cloudpickle.dumps(ray_executor_v2.RayWorkerProc) - ) - - result = patched_worker_cls().initialize_worker( - 1, - {"WORKER_ENV": "1"}, - {"DRIVER_ENV": "1"}, - assigned_physical_gpu_ids=[2, 3], - ) - - assert fp8._test_applied_configs == [config] - assert result == ( - 1, - {"WORKER_ENV": "1"}, - {"DRIVER_ENV": "1"}, - [2, 3], - ) - - def test_init_fp8_passes_modelopt_ignore_patterns_without_hf_expansion( fp8_module, monkeypatch ): @@ -1461,6 +1406,8 @@ def fake_apply( assert captured["e_score_correction_bias"] is None assert captured["routed_scaling_factor"] == 1.0 assert output.shape == x.shape + + @pytest.mark.parametrize( "use_ray_v2", ["1", "0"], ids=["ray_executor_v2", "ray_executor_v1"] ) @@ -1670,6 +1617,30 @@ class _MoERunner: ) +@GROUPED_EXPERT_KEY_SHAPES +@pytest.mark.parametrize( + "experts_dtype, expected", + [ + pytest.param(torch.float8_e4m3fn, True, id="mxfp8-middle-layer"), + pytest.param(torch.bfloat16, False, id="bf16-boundary-layer"), + ], +) +def test_is_fp8_weight_classifies_suffixless_grouped_expert_slabs( + fp8_module, + monkeypatch, + layers_prefix, + wrap_language_model, + experts_dtype, + expected, +): + fp8 = fp8_module + model = _grouped_expert_model(fp8, monkeypatch, experts_dtype, wrap_language_model) + + for suffix in ("gate_up_proj", "down_proj"): + name = f"{layers_prefix}.0.mlp.experts.{suffix}" + assert fp8._is_fp8_weight(name, model) is expected + + @GROUPED_EXPERT_KEY_SHAPES def test_load_weights_passes_grouped_experts_through_for_ignored_bf16_layers( fp8_module, monkeypatch, layers_prefix, wrap_language_model From d9246a9fda3845c6f5294f8e2ae3340b6b31537f Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 4 Sep 2026 05:11:36 -0700 Subject: [PATCH 43/76] test(mxfp8): align refit fixtures with vllm Signed-off-by: seonjinn --- .../unit/models/generation/test_vllm_fp8_quantization.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index f15268ce11c..036727baf48 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -751,7 +751,7 @@ def test_process_mxfp8_moe_refit_rejects_non_flashinfer_backend(fp8_module): with pytest.raises( NotImplementedError, - match="MXFP8 MoE refit layout conversion only supports FLASHINFER_TRTLLM", + match="requires the monolithic FlashInfer TRTLLM backend", ): fp8_module.process_weights_after_loading_mxfp8_moe(quant_method, object()) @@ -779,12 +779,14 @@ def test_process_mxfp8_moe_initializes_kernel_once(fp8_module, monkeypatch): layer.w2_weight_scale.weight_loader = object() layer._expert_routing_tables = lambda: (None, None, None) moe_config = types.SimpleNamespace(is_act_and_mul=False) - quant_config = object() - experts_cls = object() + quant_config = types.SimpleNamespace(w1_scale=None, w2_scale=None) + experts_cls = types.SimpleNamespace(is_monolithic=lambda: True) quant_config_calls = [] def get_quant_config(_layer): quant_config_calls.append(_layer) + quant_config.w1_scale = _layer.w13_weight_scale + quant_config.w2_scale = _layer.w2_weight_scale return quant_config quant_method = types.SimpleNamespace( @@ -792,6 +794,7 @@ def get_quant_config(_layer): moe_kernel=None, mxfp8_backend=Fp8MoeBackend.FLASHINFER_TRTLLM, experts_cls=experts_cls, + weight_block_size=[32, 32], get_fused_moe_quant_config=get_quant_config, ) kernel = object() From 443a9f9d80ce0e535ed80365f3d46fee5ecf2531 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 4 Sep 2026 05:21:16 -0700 Subject: [PATCH 44/76] test(mxfp8): model routed expert config Signed-off-by: seonjinn --- tests/unit/models/generation/test_vllm_fp8_quantization.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 036727baf48..3c1652e16b8 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -779,6 +779,7 @@ def test_process_mxfp8_moe_initializes_kernel_once(fp8_module, monkeypatch): layer.w2_weight_scale.weight_loader = object() layer._expert_routing_tables = lambda: (None, None, None) moe_config = types.SimpleNamespace(is_act_and_mul=False) + layer.moe_config = moe_config quant_config = types.SimpleNamespace(w1_scale=None, w2_scale=None) experts_cls = types.SimpleNamespace(is_monolithic=lambda: True) quant_config_calls = [] From b8159de8e19ff66e7592113ac0b4a4f0e703ed0b Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 4 Sep 2026 05:24:31 -0700 Subject: [PATCH 45/76] fix(vllm): reroute grouped MXFP8 scale sidecars Signed-off-by: seonjinn --- .../generation/vllm/quantization/fp8.py | 72 ++++++++++++- .../generation/test_vllm_fp8_quantization.py | 100 +++++++++++++++++- 2 files changed, 166 insertions(+), 6 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index 029025acea1..5f88ead305d 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -539,12 +539,20 @@ def _get_module_from_param_name(model, name: str): "mlp.experts.gate_up_proj", "mlp.experts.down_proj", ) +_SCALE_FROM_CHECKPOINT_SUFFIX = "_scale_from_checkpoint" def _is_grouped_expert_weight(name: str) -> bool: return name.endswith(_GROUPED_EXPERT_WEIGHT_SUFFIXES) +def _grouped_expert_weight_name_from_scale(name: str) -> str | None: + if not name.endswith(_SCALE_FROM_CHECKPOINT_SUFFIX): + return None + weight_name = name.removesuffix(_SCALE_FROM_CHECKPOINT_SUFFIX) + return weight_name if _is_grouped_expert_weight(weight_name) else None + + def _is_fp8_weight(name, model): if name not in fp8_state.seen_params: fp8_state.seen_params.add(name) @@ -592,11 +600,23 @@ def load_weights(weights, model_runner): model = model_runner.model for k, v in weights: - # Grouped MoE experts arrive as fused slabs without a ``.weight`` suffix - # (so `_is_fp8_weight` would skip them) and vLLM's grouped loader cannot - # load their per-block scales. Expand them into the per-expert FP8 (w13, w2 -> w1, w2, and w3) - # layout, then reshape to 2D [num_experts, out_features, in_features] -> [num_experts*out_features, in_features] - # so the block scales can be quantized and routed correctly. + grouped_weight_name = _grouped_expert_weight_name_from_scale(k) + if grouped_weight_name is not None: + is_prequantized_mx = ( + global_fp8_config is not None + and global_fp8_config.is_mx + and global_fp8_config.refit_prequantize + ) + if is_prequantized_mx and _is_fp8_weight(grouped_weight_name, model): + weights_quantized.extend(_reroute_grouped_moe_expert_scale(k, v)) + else: + weights_quantized.append((k, v)) + continue + # Grouped MoE experts arrive as fused slabs without a ``.weight`` suffix, + # and vLLM's grouped loader cannot load receiver-quantized per-block + # scales. Expand them into the per-expert FP8 layout so the block scales + # can be quantized and routed correctly. Prequantized MXFP8 weights stay + # fused; their scale sidecars are rerouted by the branch above. if _is_grouped_expert_weight(k): # Quantize only if vLLM built this layer's experts as FP8. Experts # covered by ``ignored_layers`` (num_{first,last}_layers_in_bf16 / @@ -859,6 +879,48 @@ def _expand_grouped_moe_expert_to_fp8(key, weight): return entries +def _reroute_grouped_moe_expert_scale( + key: str, scale: torch.Tensor +) -> list[tuple[str, torch.Tensor]]: + """Route suffixless grouped MXFP8 scales through vLLM's expert mapping. + + vLLM 0.25.1 treats every 3D expert tensor as a fused weight and applies + weight-orientation heuristics that transpose the K-compressed W13 scale. + Match the ModelOpt refit backend: emit W13 as per-expert 2D gate/up scales + and keep W2 batched behind the expert-zero checkpoint route. + """ + if scale.ndim != 3: + raise ValueError(f"Grouped MXFP8 scale {key!r} must be 3D, got {scale.ndim}D.") + + weight_name = key.removesuffix(_SCALE_FROM_CHECKPOINT_SUFFIX) + base, projection = weight_name.rsplit(".", 1) + if projection == "down_proj": + return [ + ( + f"{base}.0.down_proj.weight_scale_from_checkpoint", + scale, + ) + ] + + if scale.shape[1] % 2 != 0: + raise ValueError( + f"Grouped gate/up MXFP8 scale {key!r} must have an even projection " + f"dimension, got {tuple(scale.shape)}." + ) + gate_scale, up_scale = scale.chunk(2, dim=1) + return [ + ( + f"{base}.{expert_id}.{shard_name}.weight_scale_from_checkpoint", + expert_scale, + ) + for shard_name, grouped_scale in ( + ("gate_proj", gate_scale), + ("up_proj", up_scale), + ) + for expert_id, expert_scale in enumerate(grouped_scale.unbind(0)) + ] + + # Ref: https://github.com/vllm-project/vllm/blob/275de34170654274616082721348b7edd9741d32/vllm/model_executor/layers/quantization/utils/fp8_utils.py#L1175 # Patches this method to not create new torch.nn.Parameter for layer weights # to maintain weight loaders. diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 3c1652e16b8..2d79b002011 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -1660,26 +1660,45 @@ def test_load_weights_passes_grouped_experts_through_for_ignored_bf16_layers( import torch fp8 = fp8_module + fp8.global_fp8_config = types.SimpleNamespace( + is_mx=True, + refit_prequantize=True, + ) model = _grouped_expert_model(fp8, monkeypatch, torch.bfloat16, wrap_language_model) loaded = [] model.load_weights = lambda pairs: loaded.extend(pairs) gate_up = torch.randn(2, 256, 128).to(torch.bfloat16) + gate_up_scale = torch.ones(2, 256, 4, dtype=torch.uint8) down = torch.randn(2, 128, 128).to(torch.bfloat16) + down_scale = torch.ones(2, 128, 4, dtype=torch.uint8) fp8.load_weights( [ (f"{layers_prefix}.0.mlp.experts.gate_up_proj", gate_up), + ( + f"{layers_prefix}.0.mlp.experts.gate_up_proj_scale_from_checkpoint", + gate_up_scale, + ), (f"{layers_prefix}.0.mlp.experts.down_proj", down), + ( + f"{layers_prefix}.0.mlp.experts.down_proj_scale_from_checkpoint", + down_scale, + ), ], types.SimpleNamespace(model=model), + model_load_weights=model.load_weights, ) assert [k for k, _ in loaded] == [ f"{layers_prefix}.0.mlp.experts.gate_up_proj", + f"{layers_prefix}.0.mlp.experts.gate_up_proj_scale_from_checkpoint", f"{layers_prefix}.0.mlp.experts.down_proj", + f"{layers_prefix}.0.mlp.experts.down_proj_scale_from_checkpoint", ] assert loaded[0][1] is gate_up - assert loaded[1][1] is down + assert loaded[1][1] is gate_up_scale + assert loaded[2][1] is down + assert loaded[3][1] is down_scale # Pass-through is also what a failed lookup produces, so pin that the # bf16 RoutedExperts was actually resolved. assert isinstance( @@ -1690,6 +1709,85 @@ def test_load_weights_passes_grouped_experts_through_for_ignored_bf16_layers( ) +@GROUPED_EXPERT_KEY_SHAPES +def test_load_weights_reroutes_prequantized_grouped_expert_scale_sidecars( + fp8_module, monkeypatch, layers_prefix, wrap_language_model +): + """Avoid vLLM 0.25.1's fused-3D transpose for grouped MXFP8 scales. + + Qwen3.5's W13 sidecar is [E, 1024, 64]. The fused loader transposes it + to [E, 64, 1024], then fails copying a shard whose dim 1 is 1024 into + expert scale storage whose dim 1 is 64. + """ + import torch + + fp8 = fp8_module + fp8.global_fp8_config = types.SimpleNamespace( + use_weight_pow2_scale=False, + is_mx=True, + refit_prequantize=True, + ) + model = _grouped_expert_model( + fp8, monkeypatch, torch.float8_e4m3fn, wrap_language_model + ) + loaded = [] + + num_experts, intermediate, hidden = 2, 512, 2048 + gate_up = torch.ones( + num_experts, + 2 * intermediate, + hidden, + dtype=torch.float8_e4m3fn, + ) + gate_up_scale = torch.arange( + num_experts * 2 * intermediate * (hidden // 32), dtype=torch.uint8 + ).reshape(num_experts, 2 * intermediate, hidden // 32) + down = torch.ones( + num_experts, + hidden, + intermediate, + dtype=torch.float8_e4m3fn, + ) + down_scale = torch.arange( + num_experts * hidden * (intermediate // 32), dtype=torch.uint8 + ).reshape(num_experts, hidden, intermediate // 32) + + fp8.load_weights( + [ + (f"{layers_prefix}.0.mlp.experts.gate_up_proj", gate_up), + ( + f"{layers_prefix}.0.mlp.experts.gate_up_proj_scale_from_checkpoint", + gate_up_scale, + ), + (f"{layers_prefix}.0.mlp.experts.down_proj", down), + ( + f"{layers_prefix}.0.mlp.experts.down_proj_scale_from_checkpoint", + down_scale, + ), + ], + types.SimpleNamespace(model=model), + model_load_weights=lambda pairs: loaded.extend(pairs), + ) + + base = f"{layers_prefix}.0.mlp.experts" + assert [name for name, _ in loaded] == [ + f"{base}.gate_up_proj", + f"{base}.0.gate_proj.weight_scale_from_checkpoint", + f"{base}.1.gate_proj.weight_scale_from_checkpoint", + f"{base}.0.up_proj.weight_scale_from_checkpoint", + f"{base}.1.up_proj.weight_scale_from_checkpoint", + f"{base}.down_proj", + f"{base}.0.down_proj.weight_scale_from_checkpoint", + ] + assert loaded[0][1] is gate_up + torch.testing.assert_close(loaded[1][1], gate_up_scale[0, :intermediate]) + torch.testing.assert_close(loaded[2][1], gate_up_scale[1, :intermediate]) + torch.testing.assert_close(loaded[3][1], gate_up_scale[0, intermediate:]) + torch.testing.assert_close(loaded[4][1], gate_up_scale[1, intermediate:]) + assert loaded[5][1] is down + assert loaded[6][1] is down_scale + + def _assert_dequant_close(weight_fp8, scale_inv, source_bf16): """Dequantized FP8 must match the bf16 source within e4m3 half-ULP. From 8d38b28a4d829d591a32bcae3877086a7a60f8db Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 4 Sep 2026 05:31:09 -0700 Subject: [PATCH 46/76] test(mxfp8): patch modular kernel factory Signed-off-by: seonjinn --- .../models/generation/test_vllm_fp8_quantization.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 2d79b002011..e9b6ba4e668 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -784,10 +784,10 @@ def test_process_mxfp8_moe_initializes_kernel_once(fp8_module, monkeypatch): experts_cls = types.SimpleNamespace(is_monolithic=lambda: True) quant_config_calls = [] - def get_quant_config(_layer): - quant_config_calls.append(_layer) - quant_config.w1_scale = _layer.w13_weight_scale - quant_config.w2_scale = _layer.w2_weight_scale + def make_quant_config(**kwargs): + quant_config_calls.append(kwargs["layer"]) + quant_config.w1_scale = kwargs["w1_scale"] + quant_config.w2_scale = kwargs["w2_scale"] return quant_config quant_method = types.SimpleNamespace( @@ -796,7 +796,6 @@ def get_quant_config(_layer): mxfp8_backend=Fp8MoeBackend.FLASHINFER_TRTLLM, experts_cls=experts_cls, weight_block_size=[32, 32], - get_fused_moe_quant_config=get_quant_config, ) kernel = object() kernel_calls = [] @@ -810,7 +809,7 @@ def shuffle(*args): monkeypatch.setattr(fp8, "_shuffle_mxfp8_moe_batched", shuffle) from vllm.model_executor import parameter as vllm_parameter - from vllm.model_executor.layers.quantization import fp8 as vllm_fp8 + from vllm.model_executor.layers.fused_moe.oracle import fp8 as vllm_fp8 monkeypatch.setattr(vllm_parameter, "get_tensor_model_parallel_rank", lambda: 0) monkeypatch.setattr( @@ -821,6 +820,7 @@ def make_kernel(**kwargs): kernel_calls.append(kwargs) return kernel + monkeypatch.setattr(vllm_fp8, "make_fp8_moe_quant_config", make_quant_config) monkeypatch.setattr(vllm_fp8, "make_fp8_moe_kernel", make_kernel) fp8.process_weights_after_loading_mxfp8_moe(quant_method, layer) From 658d3852b6de96f6ed72df161e0448232c473bb7 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 4 Sep 2026 05:38:18 -0700 Subject: [PATCH 47/76] test(mxfp8): use block-aligned hidden size Signed-off-by: seonjinn --- .../models/generation/test_vllm_fp8_quantization.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index e9b6ba4e668..5c8c47df657 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -1235,11 +1235,11 @@ def fake_batched_shuffle( layer = torch.nn.Module() layer.w13_weight = torch.nn.Parameter( - torch.arange(30, dtype=torch.float32).reshape(2, 3, 5), + torch.arange(192, dtype=torch.float32).reshape(2, 3, 32), requires_grad=False, ) layer.w2_weight = torch.nn.Parameter( - torch.arange(30, dtype=torch.float32).reshape(2, 5, 3), + torch.arange(192, dtype=torch.float32).reshape(2, 32, 3), requires_grad=False, ) layer.w13_weight_scale = torch.nn.Parameter( @@ -1247,7 +1247,7 @@ def fake_batched_shuffle( requires_grad=False, ) layer.w2_weight_scale = torch.nn.Parameter( - torch.zeros(2, 5, 1, dtype=torch.uint8), + torch.zeros(2, 32, 1, dtype=torch.uint8), requires_grad=False, ) layer.w13_weight_scale_from_checkpoint = torch.nn.Parameter( @@ -1255,7 +1255,7 @@ def fake_batched_shuffle( requires_grad=False, ) layer.w2_weight_scale_from_checkpoint = torch.nn.Parameter( - torch.zeros(2, 5, 1, dtype=torch.uint8), + torch.zeros(2, 32, 1, dtype=torch.uint8), requires_grad=False, ) moe_config = types.SimpleNamespace( @@ -1289,7 +1289,7 @@ def fake_batched_shuffle( torch.testing.assert_close(layer.w13_weight, original_w13) torch.testing.assert_close(layer.w2_weight, original_w2) - assert layer.mxfp8_unpadded_hidden_size == 5 + assert layer.mxfp8_unpadded_hidden_size == 32 assert layer.mxfp8_padded_hidden_size == 512 assert layer.mxfp8_unpadded_intermediate_size_per_partition == 3 assert layer.mxfp8_padded_intermediate_size_per_partition == 128 From 172d3467ce46d4f42868f857b6fa00d26b048dfa Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 4 Sep 2026 05:53:23 -0700 Subject: [PATCH 48/76] test(mxfp8): exercise unpadded kernel reuse Signed-off-by: seonjinn --- .../models/generation/test_vllm_fp8_quantization.py | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 5c8c47df657..b5c0a37c6ab 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -767,13 +767,17 @@ def test_process_mxfp8_moe_initializes_kernel_once(fp8_module, monkeypatch): ) layer = torch.nn.Module() - layer.w13_weight = torch.nn.Parameter(torch.zeros(2, 4, 3), requires_grad=False) - layer.w2_weight = torch.nn.Parameter(torch.zeros(2, 3, 2), requires_grad=False) + layer.w13_weight = torch.nn.Parameter( + torch.zeros(2, 128, 512), requires_grad=False + ) + layer.w2_weight = torch.nn.Parameter( + torch.zeros(2, 512, 128), requires_grad=False + ) layer.w13_weight_scale = torch.nn.Parameter( - torch.zeros(2, 4, 1), requires_grad=False + torch.zeros(2, 128, 16), requires_grad=False ) layer.w2_weight_scale = torch.nn.Parameter( - torch.zeros(2, 3, 1), requires_grad=False + torch.zeros(2, 512, 4), requires_grad=False ) layer.w13_weight_scale.weight_loader = object() layer.w2_weight_scale.weight_loader = object() From 546d346078abf4247415ca99510843e9049b3ae8 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 4 Sep 2026 06:02:43 -0700 Subject: [PATCH 49/76] test(mxfp8): use current refit loader contract Signed-off-by: seonjinn --- .../generation/test_vllm_fp8_quantization.py | 35 +++++++++++++------ 1 file changed, 24 insertions(+), 11 deletions(-) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index b5c0a37c6ab..dc6d9d6c056 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -767,12 +767,8 @@ def test_process_mxfp8_moe_initializes_kernel_once(fp8_module, monkeypatch): ) layer = torch.nn.Module() - layer.w13_weight = torch.nn.Parameter( - torch.zeros(2, 128, 512), requires_grad=False - ) - layer.w2_weight = torch.nn.Parameter( - torch.zeros(2, 512, 128), requires_grad=False - ) + layer.w13_weight = torch.nn.Parameter(torch.zeros(2, 128, 512), requires_grad=False) + layer.w2_weight = torch.nn.Parameter(torch.zeros(2, 512, 128), requires_grad=False) layer.w13_weight_scale = torch.nn.Parameter( torch.zeros(2, 128, 16), requires_grad=False ) @@ -1663,6 +1659,8 @@ def test_load_weights_passes_grouped_experts_through_for_ignored_bf16_layers( """ import torch + from nemo_rl.models.generation.vllm import vllm_backend + fp8 = fp8_module fp8.global_fp8_config = types.SimpleNamespace( is_mx=True, @@ -1670,7 +1668,11 @@ def test_load_weights_passes_grouped_experts_through_for_ignored_bf16_layers( ) model = _grouped_expert_model(fp8, monkeypatch, torch.bfloat16, wrap_language_model) loaded = [] - model.load_weights = lambda pairs: loaded.extend(pairs) + monkeypatch.setattr( + vllm_backend, + "load_weights_maybe_cached", + lambda model, weights, *, cache_loader_routes: loaded.extend(weights), + ) gate_up = torch.randn(2, 256, 128).to(torch.bfloat16) gate_up_scale = torch.ones(2, 256, 4, dtype=torch.uint8) @@ -1689,8 +1691,10 @@ def test_load_weights_passes_grouped_experts_through_for_ignored_bf16_layers( down_scale, ), ], - types.SimpleNamespace(model=model), - model_load_weights=model.load_weights, + types.SimpleNamespace( + model=model, + vllm_config=types.SimpleNamespace(additional_config={}), + ), ) assert [k for k, _ in loaded] == [ @@ -1725,6 +1729,8 @@ def test_load_weights_reroutes_prequantized_grouped_expert_scale_sidecars( """ import torch + from nemo_rl.models.generation.vllm import vllm_backend + fp8 = fp8_module fp8.global_fp8_config = types.SimpleNamespace( use_weight_pow2_scale=False, @@ -1735,6 +1741,11 @@ def test_load_weights_reroutes_prequantized_grouped_expert_scale_sidecars( fp8, monkeypatch, torch.float8_e4m3fn, wrap_language_model ) loaded = [] + monkeypatch.setattr( + vllm_backend, + "load_weights_maybe_cached", + lambda model, weights, *, cache_loader_routes: loaded.extend(weights), + ) num_experts, intermediate, hidden = 2, 512, 2048 gate_up = torch.ones( @@ -1769,8 +1780,10 @@ def test_load_weights_reroutes_prequantized_grouped_expert_scale_sidecars( down_scale, ), ], - types.SimpleNamespace(model=model), - model_load_weights=lambda pairs: loaded.extend(pairs), + types.SimpleNamespace( + model=model, + vllm_config=types.SimpleNamespace(additional_config={}), + ), ) base = f"{layers_prefix}.0.mlp.experts" From db15416df244c0ad8906ec110ed854548c354600 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Fri, 4 Sep 2026 06:13:43 -0700 Subject: [PATCH 50/76] test(mxfp8): update remaining loader fixtures Signed-off-by: seonjinn --- .../generation/test_vllm_fp8_quantization.py | 18 +++++++++++++++--- 1 file changed, 15 insertions(+), 3 deletions(-) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index dc6d9d6c056..fc807da4bcb 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -1837,6 +1837,8 @@ def test_load_weights_expands_grouped_experts_for_fp8_layers( """ import torch + from nemo_rl.models.generation.vllm import vllm_backend + fp8 = fp8_module fp8.global_fp8_config = types.SimpleNamespace( use_weight_pow2_scale=False, is_mx=False @@ -1845,7 +1847,11 @@ def test_load_weights_expands_grouped_experts_for_fp8_layers( fp8, monkeypatch, torch.float8_e4m3fn, wrap_language_model ) loaded = [] - model.load_weights = lambda pairs: loaded.extend(pairs) + monkeypatch.setattr( + vllm_backend, + "load_weights_maybe_cached", + lambda model, weights, *, cache_loader_routes: loaded.extend(weights), + ) intermediate, hidden = 256, 384 gate_up = torch.randn(2, 2 * intermediate, hidden).to(torch.bfloat16) @@ -1855,7 +1861,10 @@ def test_load_weights_expands_grouped_experts_for_fp8_layers( (f"{layers_prefix}.0.mlp.experts.gate_up_proj", gate_up), (f"{layers_prefix}.0.mlp.experts.down_proj", down), ], - types.SimpleNamespace(model=model), + types.SimpleNamespace( + model=model, + vllm_config=types.SimpleNamespace(additional_config={}), + ), ) base = f"{layers_prefix}.0.mlp.experts" @@ -1904,5 +1913,8 @@ def test_load_weights_rejects_grouped_experts_for_mxfp8(fp8_module, monkeypatch) torch.randn(2, 512, 384).to(torch.bfloat16), ) ], - types.SimpleNamespace(model=model), + types.SimpleNamespace( + model=model, + vllm_config=types.SimpleNamespace(additional_config={}), + ), ) From 1d4d0134b9ef89ccfd107ef4f3688a0a4290084c Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sun, 6 Sep 2026 19:21:29 -0700 Subject: [PATCH 51/76] test(refit): reject native loader cache wrapping Signed-off-by: seonjinn --- .../generation/test_vllm_refit_loader.py | 45 +++++++++++++++++++ 1 file changed, 45 insertions(+) diff --git a/tests/unit/models/generation/test_vllm_refit_loader.py b/tests/unit/models/generation/test_vllm_refit_loader.py index 7878eac51a8..5c80bd8f690 100644 --- a/tests/unit/models/generation/test_vllm_refit_loader.py +++ b/tests/unit/models/generation/test_vllm_refit_loader.py @@ -241,6 +241,51 @@ def load_weights(self, *, weights): torch.testing.assert_close(model.default, second_default) +@pytest.mark.vllm +def test_refit_loader_cache_does_not_wrap_vllm_online_process_loader(): + from nemo_rl.models.generation.vllm.vllm_backend import ( + load_weights_maybe_cached, + ) + + loader_calls = [] + + def online_process_loader(param, loaded_weight): + loader_calls.append(loaded_weight) + with torch.no_grad(): + param.copy_(loaded_weight) + + class Model: + def __init__(self): + self.param = torch.nn.Parameter(torch.zeros(1), requires_grad=False) + self.param.weight_loader = online_process_loader + self.load_calls = 0 + + def named_parameters(self): + return [("param", self.param)] + + def load_weights(self, *, weights): + self.load_calls += 1 + assert self.param.weight_loader is online_process_loader + for name, weight in weights: + self.param.weight_loader(self.param, weight) + return {name for name, _weight in weights} + + model = Model() + + assert load_weights_maybe_cached( + model, [("param", torch.tensor([1.0]))], cache_loader_routes=True + ) == {"param"} + assert load_weights_maybe_cached( + model, [("param", torch.tensor([2.0]))], cache_loader_routes=True + ) == {"param"} + + assert model.load_calls == 2 + assert len(loader_calls) == 2 + assert model._nrl_refit_loader_cache.calls == {} + assert model._nrl_refit_loader_cache.uncached == {"param"} + torch.testing.assert_close(model.param, torch.tensor([2.0])) + + @pytest.mark.vllm def test_refit_loader_cache_invalidates_replaced_parameter(monkeypatch): from nemo_rl.models.generation.vllm.vllm_backend import ( From 1b5e6420a41aa5a9204f7a7eb728e791c9425481 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sun, 6 Sep 2026 20:09:10 -0700 Subject: [PATCH 52/76] fix(refit): preserve vLLM online loaders Signed-off-by: seonjinn --- nemo_rl/models/generation/vllm/vllm_backend.py | 14 +++++++++----- 1 file changed, 9 insertions(+), 5 deletions(-) diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index 7b2ff013cb7..36d329f4b5d 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -263,11 +263,15 @@ def recorder(param, loaded_weight, *args, **kwargs): try: for param_name, param in model.named_parameters(): loader = getattr(param, "weight_loader", None) - # Leave default_weight_loader params unwrapped: model - # load_weights implementations dispatch on - # `weight_loader == default_weight_loader` with a different - # argument list. - if loader is None or loader is default_weight_loader: + # Leave loaders owned by vLLM's loading lifecycle unwrapped. + # The default loader may use a different argument list, while the + # online loader must remain installed so vLLM can finalize each + # layer without replaying through this cache recorder. + if ( + loader is None + or loader is default_weight_loader + or getattr(loader, "__name__", None) == "online_process_loader" + ): continue cache.snapshot[param_name] = param originals.append((param, loader)) From b7d4ce6bff753dc765311f05e89c92efc4a2c68c Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sun, 6 Sep 2026 20:50:48 -0700 Subject: [PATCH 53/76] test(refit): avoid duplicate kv scale finalization Signed-off-by: seonjinn --- tests/unit/models/generation/test_vllm_backend.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/unit/models/generation/test_vllm_backend.py b/tests/unit/models/generation/test_vllm_backend.py index b5b4ffa5dd7..a1b3381c282 100644 --- a/tests/unit/models/generation/test_vllm_backend.py +++ b/tests/unit/models/generation/test_vllm_backend.py @@ -374,7 +374,7 @@ def test_unquantized_nccl_reshard_keeps_existing_refit_lifecycle(monkeypatch): process.assert_called_once_with(model, model_config, ext.device) ext._maybe_process_mtp_drafter_after_loading.assert_called_once_with() - ext._maybe_process_fp8_kv_cache.assert_called_once_with() + ext._maybe_process_fp8_kv_cache.assert_not_called() @pytest.mark.vllm From 201ad9efc69d1a053b9d62230d53e71226f71c88 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sun, 6 Sep 2026 20:57:59 -0700 Subject: [PATCH 54/76] test(refit): avoid duplicate fp8 kv finalization Signed-off-by: seonjinn --- tests/unit/models/generation/test_vllm_backend.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/unit/models/generation/test_vllm_backend.py b/tests/unit/models/generation/test_vllm_backend.py index a1b3381c282..6eda2bf5db5 100644 --- a/tests/unit/models/generation/test_vllm_backend.py +++ b/tests/unit/models/generation/test_vllm_backend.py @@ -526,7 +526,7 @@ def test_fp8_flashinfer_trtllm_keeps_existing_refit_lifecycle(monkeypatch): process.assert_called_once_with(model, model_config, ext.device) ext._maybe_process_mtp_drafter_after_loading.assert_called_once_with() - ext._maybe_process_fp8_kv_cache.assert_called_once_with() + ext._maybe_process_fp8_kv_cache.assert_not_called() @pytest.mark.vllm From 74381536a779420d7e9767cd9c42694953687801 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sun, 6 Sep 2026 21:04:47 -0700 Subject: [PATCH 55/76] test(refit): use a realized module fixture Signed-off-by: seonjinn --- tests/unit/models/generation/test_vllm_backend.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/unit/models/generation/test_vllm_backend.py b/tests/unit/models/generation/test_vllm_backend.py index 6eda2bf5db5..7b850c26c61 100644 --- a/tests/unit/models/generation/test_vllm_backend.py +++ b/tests/unit/models/generation/test_vllm_backend.py @@ -763,7 +763,7 @@ def test_prepare_refit_info_reports_only_fp8_weights(monkeypatch, enabled): ext = vllm_backend.VllmInternalWorkerExtension.__new__( vllm_backend.VllmInternalWorkerExtension ) - model = object() + model = torch.nn.Module() config = object() ext.model_runner = SimpleNamespace(model=model, vllm_config=config) state_dict_info = { From 08e7d10839bfccd7b80f397bd5e3a2f792ff76ef Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sun, 6 Sep 2026 21:12:55 -0700 Subject: [PATCH 56/76] test(refit): expect serialized FP8 config Signed-off-by: seonjinn --- tests/unit/models/generation/test_vllm_backend.py | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/tests/unit/models/generation/test_vllm_backend.py b/tests/unit/models/generation/test_vllm_backend.py index 7b850c26c61..fb0b0000e55 100644 --- a/tests/unit/models/generation/test_vllm_backend.py +++ b/tests/unit/models/generation/test_vllm_backend.py @@ -1253,17 +1253,20 @@ async def test_async_weight_update_fails_when_encoder_cache_reset_fails(): @pytest.mark.vllm -def test_worker_prepare_refit_info_forwards_state_dict_info(): +def test_worker_prepare_refit_info_forwards_state_dict_info(monkeypatch): + from nemo_rl.models.generation.vllm.quantization import fp8 from nemo_rl.models.generation.vllm.vllm_worker import VllmGenerationWorkerImpl worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) worker.llm = SimpleNamespace(collective_rpc=MagicMock()) state_dict_info = {"model.weight": object()} + serialized_config = {"is_mx": True} + monkeypatch.setattr(fp8, "serialize_fp8_config", lambda: serialized_config) worker.prepare_refit_info(state_dict_info) assert worker.llm.collective_rpc.call_args_list == [ - call("prepare_refit_info", args=(state_dict_info,)), + call("prepare_refit_info", args=(state_dict_info, serialized_config)), ] @@ -1328,7 +1331,8 @@ def test_generation_prepare_refit_info_allows_prequantized_mxfp8_grouped_moe( @pytest.mark.vllm @pytest.mark.asyncio -async def test_async_worker_prepare_refit_info_forwards_state_dict_info(): +async def test_async_worker_prepare_refit_info_forwards_state_dict_info(monkeypatch): + from nemo_rl.models.generation.vllm.quantization import fp8 from nemo_rl.models.generation.vllm.vllm_worker_async import ( VllmAsyncGenerationWorkerImpl, ) @@ -1336,11 +1340,13 @@ async def test_async_worker_prepare_refit_info_forwards_state_dict_info(): worker = VllmAsyncGenerationWorkerImpl.__new__(VllmAsyncGenerationWorkerImpl) worker.llm = SimpleNamespace(collective_rpc=AsyncMock()) state_dict_info = {"model.weight": object()} + serialized_config = {"is_mx": True} + monkeypatch.setattr(fp8, "serialize_fp8_config", lambda: serialized_config) await worker.prepare_refit_info_async(state_dict_info) assert worker.llm.collective_rpc.await_args_list == [ - call("prepare_refit_info", args=(state_dict_info,)), + call("prepare_refit_info", args=(state_dict_info, serialized_config)), ] From d0034830d7f7b98a85cf3834cc3096b285525174 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sun, 6 Sep 2026 22:23:38 -0700 Subject: [PATCH 57/76] test(refit): follow collective sender contract Signed-off-by: seonjinn --- .../weight_sync/test_weight_synchronizer.py | 26 ------------------- 1 file changed, 26 deletions(-) diff --git a/tests/unit/weight_sync/test_weight_synchronizer.py b/tests/unit/weight_sync/test_weight_synchronizer.py index a320742950d..ae05d47c24b 100644 --- a/tests/unit/weight_sync/test_weight_synchronizer.py +++ b/tests/unit/weight_sync/test_weight_synchronizer.py @@ -643,30 +643,6 @@ def test_sync_weights_passes_kv_scales(self, mock_ray): call_kwargs = policy.broadcast_weights_for_collective.call_args assert call_kwargs.kwargs["kv_scales"] == kv_scales - @patch("nemo_rl.weight_sync.collective_weight_synchronizer.ray") - def test_sync_weights_forwards_fixed_buffer_size(self, mock_ray): - mock_ray.get.return_value = [True] - policy = _mock_policy() - gen = _mock_generation() - sync = CollectiveWeightSynchronizer( - policy, - gen, - _mock_cluster(), - _mock_cluster(), - refit_buffer_size_gb=1.5, - ) - - sync.sync_weights() - - expected_bytes = int(1.5 * 1024**3) - policy.broadcast_weights_for_collective.assert_called_once_with( - kv_scales=None, - buffer_size_bytes=expected_bytes, - ) - gen.update_weights_from_collective.assert_called_once_with( - buffer_size_bytes=expected_bytes - ) - @patch("nemo_rl.weight_sync.collective_weight_synchronizer.ray") def test_sync_weights_raises_on_failure(self, mock_ray): mock_ray.get.side_effect = [ @@ -1014,10 +990,8 @@ def test_non_colocated_vllm_returns_collective(self): colocated=False, train_cluster=_mock_cluster(), inference_cluster=_mock_cluster(), - refit_buffer_size_gb=1.5, ) assert isinstance(sync, CollectiveWeightSynchronizer) - assert sync._buffer_size_bytes == int(1.5 * 1024**3) def test_non_colocated_dynamo_returns_collective(self): sync = create_weight_synchronizer( From 3ace72848230c8b716dcd6032d12d5f4e766a868 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Mon, 24 Aug 2026 14:08:58 -0700 Subject: [PATCH 58/76] perf(refit): batch MXFP8 expert prequantization Signed-off-by: seonjinn --- .../vllm/quantization/fp8_train_utils.py | 145 ++++++++++++++++++ .../policy/workers/megatron_policy_worker.py | 19 ++- .../models/generation/test_mxfp8_prequant.py | 126 +++++++++++++++ .../models/policy/test_megatron_worker.py | 39 +++++ 4 files changed, 327 insertions(+), 2 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py index c4d2df7e290..a01aa118358 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py @@ -13,11 +13,19 @@ # limitations under the License. +import re +from collections.abc import Callable, Iterable, Iterator + import torch MXFP8_BLOCK_SIZE = 32 MXFP8_VALUE_DTYPE = torch.float8_e4m3fn +_EXPERT_WEIGHT_PATTERN = re.compile( + r"^(?P.+\.experts)\.(?P\d+)\." + r"(?Pgate_proj|up_proj|down_proj)\.weight$" +) + def _mxfp8_e4m3_quantize_torch( x: torch.Tensor, @@ -95,6 +103,143 @@ def mxfp8_e4m3_quantize_for_refit( return x_q, x_scales +def iter_mxfp8_prequantized_params( + params: Iterable[tuple[str, torch.Tensor]], + selected_names: set[str], + *, + quantize_fn: Callable[ + [torch.Tensor], tuple[torch.Tensor, torch.Tensor] + ] = mxfp8_e4m3_quantize_for_refit, + scratch_cache: dict[tuple[torch.device, torch.dtype], torch.Tensor] | None = None, + max_experts_per_batch: int = 16, +) -> Iterator[tuple[str, torch.Tensor]]: + """Batch expert weights while preserving the existing refit wire entries. + + Selected MoE experts with matching layer, projection, shape, dtype, and + device are quantized together in bounded chunks. Other selected weights use + the existing per-tensor path, and unselected weights pass through unchanged. + + Args: + params: Exported Hugging Face parameter names and tensors. + selected_names: Parameter names selected for MXFP8 prequantization. + quantize_fn: MXFP8 quantization function. + scratch_cache: Reusable stacking buffers keyed by device and dtype. + max_experts_per_batch: Maximum number of experts per quantization call. + + Yields: + Weight and scale entries accepted by the vLLM refit receiver. + """ + if max_experts_per_batch <= 0: + raise ValueError("max_experts_per_batch must be positive") + if scratch_cache is None: + scratch_cache = {} + + pending: dict[tuple[str, str], list[tuple[int, str, torch.Tensor]]] = {} + current_prefix: str | None = None + + def quantize_one(name: str, tensor: torch.Tensor) -> list[tuple[str, torch.Tensor]]: + if tensor.dtype == torch.float8_e4m3fn: + raise ValueError( + "MXFP8 prequantization requires BF16 trainer-exported weights; " + f"{name} is already stored as E4M3." + ) + value, scale = quantize_fn(tensor) + return [(name, value), (name + "_scale_from_checkpoint", scale)] + + def flush_group( + group_key: tuple[str, str], + ) -> Iterator[tuple[str, torch.Tensor]]: + group = pending.pop(group_key) + group.sort(key=lambda item: item[0]) + while group: + chunk = group[:max_experts_per_batch] + del group[:max_experts_per_batch] + tensors = [tensor for _expert_id, _name, tensor in chunk] + batchable = len(chunk) > 1 and len({item[0] for item in chunk}) == len( + chunk + ) + if batchable: + first = tensors[0] + batchable = all( + tensor.shape == first.shape + and tensor.dtype == first.dtype + and tensor.device == first.device + and tensor.layout is torch.strided + for tensor in tensors + ) + if not batchable: + for _expert_id, name, tensor in chunk: + yield from quantize_one(name, tensor) + continue + + first = tensors[0] + if first.dtype == torch.float8_e4m3fn: + raise ValueError( + "MXFP8 prequantization requires BF16 trainer-exported weights." + ) + required_numel = len(chunk) * first.numel() + cache_key = (first.device, first.dtype) + scratch = scratch_cache.get(cache_key) + if scratch is None or scratch.numel() < required_numel: + scratch = torch.empty( + required_numel, + dtype=first.dtype, + device=first.device, + ) + scratch_cache[cache_key] = scratch + stacked = scratch[:required_numel].view(len(chunk), *first.shape) + torch.stack(tensors, dim=0, out=stacked) + + value, scale = quantize_fn(stacked.view(-1, stacked.shape[-1])) + value = value.view_as(stacked) + scale = scale.view( + len(chunk), + *first.shape[:-1], + first.shape[-1] // MXFP8_BLOCK_SIZE, + ) + for index, (_expert_id, name, _tensor) in enumerate(chunk): + yield name, value[index] + yield name + "_scale_from_checkpoint", scale[index] + + def flush_pending() -> Iterator[tuple[str, torch.Tensor]]: + while pending: + yield from flush_group(next(iter(pending))) + + for name, tensor in params: + match = _EXPERT_WEIGHT_PATTERN.match(name) if name in selected_names else None + if match is None: + if pending: + yield from flush_pending() + current_prefix = None + if name in selected_names: + yield from quantize_one(name, tensor) + else: + yield name, tensor + continue + + prefix = match.group("prefix") + projection = match.group("projection") + if current_prefix is not None and prefix != current_prefix: + yield from flush_pending() + elif projection == "down_proj" and any( + key[1] != "down_proj" for key in pending + ): + yield from flush_pending() + elif projection != "down_proj" and any( + key[1] == "down_proj" for key in pending + ): + yield from flush_pending() + current_prefix = prefix + group_key = (prefix, projection) + group = pending.setdefault(group_key, []) + group.append((int(match.group("expert_id")), name, tensor)) + if len(group) == max_experts_per_batch: + yield from flush_group(group_key) + + if pending: + yield from flush_pending() + + def get_vllm_qkv_scale_names(layer_idx: int) -> dict[str, str]: """Get vLLM-compatible parameter names for Q/K/V FP8 scales. diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index e4359572599..e701fdc5a86 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -485,6 +485,9 @@ def __init__( self._refit_param_info_hf: Optional[ dict[str, tuple[torch.Size, torch.dtype]] ] = None + self._mxfp8_prequant_scratch_cache: dict[ + tuple[torch.device, torch.dtype], torch.Tensor + ] = {} # Pinned host staging for the reference-policy swap; only populated when # megatron_cfg["pinned_reference_swap"] is enabled. Buffer contents are # only live within a single use_reference_model call (every copy @@ -2653,8 +2656,20 @@ def _iter_params_with_optional_kv_scales( # Yield the original parameters first, MXFP8-quantizing on the trainer # when pre-quantized refit is enabled for the parameter. - for name, tensor in base_iter: - yield from self._maybe_prequantize_param(name, tensor) + if self._refit_prequant_names: + # Trainer workers only need this optional vLLM helper during MXFP8 refit. + from nemo_rl.models.generation.vllm.quantization.fp8_train_utils import ( + iter_mxfp8_prequantized_params, + ) + + yield from iter_mxfp8_prequantized_params( + base_iter, + self._refit_prequant_names, + scratch_cache=self._mxfp8_prequant_scratch_cache, + ) + else: + for name, tensor in base_iter: + yield from self._maybe_prequantize_param(name, tensor) if include_draft and self.draft_model is not None: from nemo_rl.models.megatron.draft import export_eagle_weights_to_hf diff --git a/tests/unit/models/generation/test_mxfp8_prequant.py b/tests/unit/models/generation/test_mxfp8_prequant.py index 5c762b8ca46..b1d6f8af256 100644 --- a/tests/unit/models/generation/test_mxfp8_prequant.py +++ b/tests/unit/models/generation/test_mxfp8_prequant.py @@ -126,6 +126,132 @@ def test_refit_quantize_matches_receiver_path(): assert torch.equal(got_scale.reshape(-1), ref_scale.reshape(-1)) +def test_batched_expert_prequantization_preserves_wire_entries_and_reuses_scratch(): + from nemo_rl.models.generation.vllm.quantization import fp8_train_utils + + calls = [] + + def quantize(tensor): + calls.append(tuple(tensor.shape)) + scales = torch.ones( + (*tensor.shape[:-1], tensor.shape[-1] // MXFP8_BLOCK_SIZE), + dtype=torch.uint8, + ) + return tensor.clone(), scales + + def expert_name(expert_id, projection): + return f"model.layers.0.mlp.experts.{expert_id}.{projection}_proj.weight" + + params = [("model.layers.0.input_layernorm.weight", torch.ones(64))] + expected = {} + for expert_id in range(2): + for projection in ("gate", "up"): + name = expert_name(expert_id, projection) + tensor = torch.full((2, 64), expert_id + (1 if projection == "gate" else 3)) + params.append((name, tensor)) + expected[name] = tensor + for expert_id in range(2): + name = expert_name(expert_id, "down") + tensor = torch.full((4, 32), expert_id + 5) + params.append((name, tensor)) + expected[name] = tensor + + selected_names = set(expected) + scratch_cache = {} + output = dict( + fp8_train_utils.iter_mxfp8_prequantized_params( + iter(params), + selected_names, + quantize_fn=quantize, + scratch_cache=scratch_cache, + ) + ) + + assert calls == [(4, 64), (4, 64), (8, 32)] + assert output[params[0][0]] is params[0][1] + for name, tensor in expected.items(): + torch.testing.assert_close(output[name], tensor) + scale_name = name + "_scale_from_checkpoint" + assert output[scale_name].shape == (*tensor.shape[:-1], tensor.shape[-1] // 32) + assert torch.all(output[scale_name] == 1) + + scratch = next(iter(scratch_cache.values())) + first_scratch_ptr = scratch.data_ptr() + calls.clear() + list( + fp8_train_utils.iter_mxfp8_prequantized_params( + iter(params), + selected_names, + quantize_fn=quantize, + scratch_cache=scratch_cache, + ) + ) + assert calls == [(4, 64), (4, 64), (8, 32)] + assert next(iter(scratch_cache.values())).data_ptr() == first_scratch_ptr + + +def test_batched_expert_prequantization_bounds_batch_and_has_stable_order(): + from nemo_rl.models.generation.vllm.quantization import fp8_train_utils + + calls = [] + + def quantize(tensor): + calls.append(tuple(tensor.shape)) + scales = torch.ones( + (*tensor.shape[:-1], tensor.shape[-1] // MXFP8_BLOCK_SIZE), + dtype=torch.uint8, + ) + return tensor.clone(), scales + + def expert_name(expert_id, projection): + return f"model.layers.0.mlp.experts.{expert_id}.{projection}_proj.weight" + + params = [] + for expert_id in range(5): + for projection in ("gate", "up"): + params.append((expert_name(expert_id, projection), torch.ones(2, 64))) + for expert_id in range(5): + params.append((expert_name(expert_id, "down"), torch.ones(4, 32))) + + output = list( + fp8_train_utils.iter_mxfp8_prequantized_params( + iter(params), + {name for name, _tensor in params}, + quantize_fn=quantize, + max_experts_per_batch=2, + ) + ) + + expected_names = [] + for expert_ids, projection in ( + ((0, 1), "gate"), + ((0, 1), "up"), + ((2, 3), "gate"), + ((2, 3), "up"), + ((4,), "gate"), + ((4,), "up"), + ((0, 1), "down"), + ((2, 3), "down"), + ((4,), "down"), + ): + for expert_id in expert_ids: + name = expert_name(expert_id, projection) + expected_names.extend((name, name + "_scale_from_checkpoint")) + + assert [name for name, _tensor in output] == expected_names + assert calls == [ + (4, 64), + (4, 64), + (4, 64), + (4, 64), + (2, 64), + (2, 64), + (8, 32), + (8, 32), + (4, 32), + ] + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_refit_quantize_matches_receiver_quantize_mxfp8_weight(): """Sender prequantization and the receiver helper must agree bit-for-bit. diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index 10c4e552caa..24aa1ed703a 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -788,6 +788,45 @@ def test_maybe_prequantize_param_rejects_fp8_trainer_storage(): list(worker._maybe_prequantize_param(name, tensor)) +def test_iter_params_batches_expert_prequantization_and_reuses_scratch( + monkeypatch, +): + from nemo_rl.models.generation.vllm.quantization import fp8_train_utils + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + + name = "model.layers.0.mlp.experts.0.gate_proj.weight" + weight = torch.ones(2, 32, dtype=torch.bfloat16) + calls = [] + + def iter_batched(params, selected_names, *, scratch_cache): + calls.append((list(params), selected_names, scratch_cache)) + yield "batched.weight", weight + + monkeypatch.setattr(fp8_train_utils, "iter_mxfp8_prequantized_params", iter_batched) + worker = object.__new__(MegatronPolicyWorkerImpl) + worker._refit_prequant_names = {name} + worker._mxfp8_prequant_scratch_cache = {} + worker.model = object() + worker.draft_model = None + worker.refit_conversion_tasks = [] + worker.cfg = {"megatron_cfg": {"enabled": True}} + worker.megatron_bridge = SimpleNamespace( + export_hf_weights=lambda *_args, **_kwargs: iter([(name, weight)]) + ) + + first = list(worker._iter_params_with_optional_kv_scales()) + second = list(worker._iter_params_with_optional_kv_scales()) + + assert first == [("batched.weight", weight)] + assert second == first + assert len(calls) == 2 + assert calls[0][0] == [(name, weight)] + assert calls[0][1] == {name} + assert calls[0][2] is calls[1][2] + + def test_enable_refit_prequantize_rejects_blockwise_fp8_storage(): from nemo_rl.models.policy.workers.megatron_policy_worker import ( MegatronPolicyWorkerImpl, From d1848119201424445701d20a72caed3832ed1bfe Mon Sep 17 00:00:00 2001 From: seonjinn Date: Mon, 24 Aug 2026 14:15:01 -0700 Subject: [PATCH 59/76] fix(refit): support grad-enabled expert views Signed-off-by: seonjinn --- .../generation/vllm/quantization/fp8_train_utils.py | 3 ++- tests/unit/models/generation/test_mxfp8_prequant.py | 10 +++++++--- 2 files changed, 9 insertions(+), 4 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py index a01aa118358..0e2dabaa6cb 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py @@ -188,7 +188,8 @@ def flush_group( ) scratch_cache[cache_key] = scratch stacked = scratch[:required_numel].view(len(chunk), *first.shape) - torch.stack(tensors, dim=0, out=stacked) + with torch.no_grad(): + torch.stack(tensors, dim=0, out=stacked) value, scale = quantize_fn(stacked.view(-1, stacked.shape[-1])) value = value.view_as(stacked) diff --git a/tests/unit/models/generation/test_mxfp8_prequant.py b/tests/unit/models/generation/test_mxfp8_prequant.py index b1d6f8af256..3f5af39aea9 100644 --- a/tests/unit/models/generation/test_mxfp8_prequant.py +++ b/tests/unit/models/generation/test_mxfp8_prequant.py @@ -137,7 +137,7 @@ def quantize(tensor): (*tensor.shape[:-1], tensor.shape[-1] // MXFP8_BLOCK_SIZE), dtype=torch.uint8, ) - return tensor.clone(), scales + return tensor.detach().clone(), scales def expert_name(expert_id, projection): return f"model.layers.0.mlp.experts.{expert_id}.{projection}_proj.weight" @@ -147,12 +147,16 @@ def expert_name(expert_id, projection): for expert_id in range(2): for projection in ("gate", "up"): name = expert_name(expert_id, projection) - tensor = torch.full((2, 64), expert_id + (1 if projection == "gate" else 3)) + tensor = torch.full( + (2, 64), + expert_id + (1 if projection == "gate" else 3), + requires_grad=True, + ) params.append((name, tensor)) expected[name] = tensor for expert_id in range(2): name = expert_name(expert_id, "down") - tensor = torch.full((4, 32), expert_id + 5) + tensor = torch.full((4, 32), expert_id + 5, requires_grad=True) params.append((name, tensor)) expected[name] = tensor From ccd405495241404607e92d5dec224c97c2dd1fb8 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Mon, 24 Aug 2026 14:21:14 -0700 Subject: [PATCH 60/76] fix(refit): release MXFP8 scratch after export Signed-off-by: seonjinn --- .../models/policy/workers/megatron_policy_worker.py | 4 ---- tests/unit/models/policy/test_megatron_worker.py | 10 +++------- 2 files changed, 3 insertions(+), 11 deletions(-) diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index e701fdc5a86..f6e21ae5ffb 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -485,9 +485,6 @@ def __init__( self._refit_param_info_hf: Optional[ dict[str, tuple[torch.Size, torch.dtype]] ] = None - self._mxfp8_prequant_scratch_cache: dict[ - tuple[torch.device, torch.dtype], torch.Tensor - ] = {} # Pinned host staging for the reference-policy swap; only populated when # megatron_cfg["pinned_reference_swap"] is enabled. Buffer contents are # only live within a single use_reference_model call (every copy @@ -2665,7 +2662,6 @@ def _iter_params_with_optional_kv_scales( yield from iter_mxfp8_prequantized_params( base_iter, self._refit_prequant_names, - scratch_cache=self._mxfp8_prequant_scratch_cache, ) else: for name, tensor in base_iter: diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index 24aa1ed703a..c3f9aaa7b07 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -788,9 +788,7 @@ def test_maybe_prequantize_param_rejects_fp8_trainer_storage(): list(worker._maybe_prequantize_param(name, tensor)) -def test_iter_params_batches_expert_prequantization_and_reuses_scratch( - monkeypatch, -): +def test_iter_params_batches_expert_prequantization(monkeypatch): from nemo_rl.models.generation.vllm.quantization import fp8_train_utils from nemo_rl.models.policy.workers.megatron_policy_worker import ( MegatronPolicyWorkerImpl, @@ -800,14 +798,13 @@ def test_iter_params_batches_expert_prequantization_and_reuses_scratch( weight = torch.ones(2, 32, dtype=torch.bfloat16) calls = [] - def iter_batched(params, selected_names, *, scratch_cache): - calls.append((list(params), selected_names, scratch_cache)) + def iter_batched(params, selected_names): + calls.append((list(params), selected_names)) yield "batched.weight", weight monkeypatch.setattr(fp8_train_utils, "iter_mxfp8_prequantized_params", iter_batched) worker = object.__new__(MegatronPolicyWorkerImpl) worker._refit_prequant_names = {name} - worker._mxfp8_prequant_scratch_cache = {} worker.model = object() worker.draft_model = None worker.refit_conversion_tasks = [] @@ -824,7 +821,6 @@ def iter_batched(params, selected_names, *, scratch_cache): assert len(calls) == 2 assert calls[0][0] == [(name, weight)] assert calls[0][1] == {name} - assert calls[0][2] is calls[1][2] def test_enable_refit_prequantize_rejects_blockwise_fp8_storage(): From ed8b37d688c0555cbe1726341b1fed2103fed52c Mon Sep 17 00:00:00 2001 From: seonjinn Date: Mon, 24 Aug 2026 14:27:30 -0700 Subject: [PATCH 61/76] fix(refit): preserve expert export and scale contracts Signed-off-by: seonjinn --- .../vllm/quantization/fp8_train_utils.py | 31 ++--- .../models/generation/test_mxfp8_prequant.py | 107 ++++++++++++++---- 2 files changed, 103 insertions(+), 35 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py index 0e2dabaa6cb..1c26e58312a 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py @@ -110,7 +110,8 @@ def iter_mxfp8_prequantized_params( quantize_fn: Callable[ [torch.Tensor], tuple[torch.Tensor, torch.Tensor] ] = mxfp8_e4m3_quantize_for_refit, - scratch_cache: dict[tuple[torch.device, torch.dtype], torch.Tensor] | None = None, + scratch_cache: dict[tuple[torch.device, torch.dtype, int | None], torch.Tensor] + | None = None, max_experts_per_batch: int = 16, ) -> Iterator[tuple[str, torch.Tensor]]: """Batch expert weights while preserving the existing refit wire entries. @@ -123,7 +124,8 @@ def iter_mxfp8_prequantized_params( params: Exported Hugging Face parameter names and tensors. selected_names: Parameter names selected for MXFP8 prequantization. quantize_fn: MXFP8 quantization function. - scratch_cache: Reusable stacking buffers keyed by device and dtype. + scratch_cache: Reusable stacking buffers keyed by device, dtype, and + CUDA stream. max_experts_per_batch: Maximum number of experts per quantization call. Yields: @@ -178,7 +180,12 @@ def flush_group( "MXFP8 prequantization requires BF16 trainer-exported weights." ) required_numel = len(chunk) * first.numel() - cache_key = (first.device, first.dtype) + stream_id = ( + int(torch.cuda.current_stream(first.device).cuda_stream) + if first.is_cuda + else None + ) + cache_key = (first.device, first.dtype, stream_id) scratch = scratch_cache.get(cache_key) if scratch is None or scratch.numel() < required_numel: scratch = torch.empty( @@ -193,11 +200,13 @@ def flush_group( value, scale = quantize_fn(stacked.view(-1, stacked.shape[-1])) value = value.view_as(stacked) - scale = scale.view( - len(chunk), - *first.shape[:-1], - first.shape[-1] // MXFP8_BLOCK_SIZE, + scale_columns = first.shape[-1] // MXFP8_BLOCK_SIZE + scale_shape = ( + first.shape[:-1] + if scale_columns == 1 + else (*first.shape[:-1], scale_columns) ) + scale = scale.view(len(chunk), *scale_shape) for index, (_expert_id, name, _tensor) in enumerate(chunk): yield name, value[index] yield name + "_scale_from_checkpoint", scale[index] @@ -222,14 +231,6 @@ def flush_pending() -> Iterator[tuple[str, torch.Tensor]]: projection = match.group("projection") if current_prefix is not None and prefix != current_prefix: yield from flush_pending() - elif projection == "down_proj" and any( - key[1] != "down_proj" for key in pending - ): - yield from flush_pending() - elif projection != "down_proj" and any( - key[1] == "down_proj" for key in pending - ): - yield from flush_pending() current_prefix = prefix group_key = (prefix, projection) group = pending.setdefault(group_key, []) diff --git a/tests/unit/models/generation/test_mxfp8_prequant.py b/tests/unit/models/generation/test_mxfp8_prequant.py index 3f5af39aea9..e12c42e59ea 100644 --- a/tests/unit/models/generation/test_mxfp8_prequant.py +++ b/tests/unit/models/generation/test_mxfp8_prequant.py @@ -145,20 +145,17 @@ def expert_name(expert_id, projection): params = [("model.layers.0.input_layernorm.weight", torch.ones(64))] expected = {} for expert_id in range(2): - for projection in ("gate", "up"): + for projection in ("gate", "up", "down"): name = expert_name(expert_id, projection) - tensor = torch.full( - (2, 64), - expert_id + (1 if projection == "gate" else 3), - requires_grad=True, - ) + if projection == "down": + shape = (4, 32) + fill_value = expert_id + 5 + else: + shape = (2, 64) + fill_value = expert_id + (1 if projection == "gate" else 3) + tensor = torch.full(shape, fill_value, requires_grad=True) params.append((name, tensor)) expected[name] = tensor - for expert_id in range(2): - name = expert_name(expert_id, "down") - tensor = torch.full((4, 32), expert_id + 5, requires_grad=True) - params.append((name, tensor)) - expected[name] = tensor selected_names = set(expected) scratch_cache = {} @@ -176,7 +173,13 @@ def expert_name(expert_id, projection): for name, tensor in expected.items(): torch.testing.assert_close(output[name], tensor) scale_name = name + "_scale_from_checkpoint" - assert output[scale_name].shape == (*tensor.shape[:-1], tensor.shape[-1] // 32) + scale_columns = tensor.shape[-1] // MXFP8_BLOCK_SIZE + expected_scale_shape = ( + tensor.shape[:-1] + if scale_columns == 1 + else (*tensor.shape[:-1], scale_columns) + ) + assert output[scale_name].shape == expected_scale_shape assert torch.all(output[scale_name] == 1) scratch = next(iter(scratch_cache.values())) @@ -212,10 +215,9 @@ def expert_name(expert_id, projection): params = [] for expert_id in range(5): - for projection in ("gate", "up"): - params.append((expert_name(expert_id, projection), torch.ones(2, 64))) - for expert_id in range(5): - params.append((expert_name(expert_id, "down"), torch.ones(4, 32))) + for projection in ("gate", "up", "down"): + shape = (4, 32) if projection == "down" else (2, 64) + params.append((expert_name(expert_id, projection), torch.ones(*shape))) output = list( fp8_train_utils.iter_mxfp8_prequantized_params( @@ -230,12 +232,12 @@ def expert_name(expert_id, projection): for expert_ids, projection in ( ((0, 1), "gate"), ((0, 1), "up"), + ((0, 1), "down"), ((2, 3), "gate"), ((2, 3), "up"), + ((2, 3), "down"), ((4,), "gate"), ((4,), "up"), - ((0, 1), "down"), - ((2, 3), "down"), ((4,), "down"), ): for expert_id in expert_ids: @@ -246,16 +248,81 @@ def expert_name(expert_id, projection): assert calls == [ (4, 64), (4, 64), + (8, 32), (4, 64), (4, 64), + (8, 32), (2, 64), (2, 64), - (8, 32), - (8, 32), (4, 32), ] +def test_batched_expert_prequantization_matches_per_tensor_quantization(): + from nemo_rl.models.generation.vllm.quantization import fp8_train_utils + + torch.manual_seed(0) + + def expert_name(expert_id, projection): + return f"model.layers.0.mlp.experts.{expert_id}.{projection}_proj.weight" + + params = [] + for expert_id in range(3): + params.extend( + [ + (expert_name(expert_id, "gate"), torch.randn(2, 64)), + (expert_name(expert_id, "up"), torch.randn(2, 64)), + (expert_name(expert_id, "down"), torch.randn(4, 32)), + ] + ) + + output = dict( + fp8_train_utils.iter_mxfp8_prequantized_params( + params, + {name for name, _tensor in params}, + ) + ) + + for name, tensor in params: + expected_value, expected_scale = fp8_train_utils.mxfp8_e4m3_quantize_for_refit( + tensor + ) + assert torch.equal( + output[name].view(torch.uint8), expected_value.view(torch.uint8) + ) + assert torch.equal(output[name + "_scale_from_checkpoint"], expected_scale) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_batched_expert_prequantization_uses_stream_local_scratch(): + from nemo_rl.models.generation.vllm.quantization import fp8_train_utils + + def expert_name(expert_id): + return f"model.layers.0.mlp.experts.{expert_id}.gate_proj.weight" + + params = [(expert_name(i), torch.ones(2, 64, device="cuda")) for i in range(4)] + scratch_cache = {} + output = fp8_train_utils.iter_mxfp8_prequantized_params( + params, + {name for name, _tensor in params}, + quantize_fn=lambda tensor: ( + tensor.clone(), + torch.ones((*tensor.shape[:-1], 2), dtype=torch.uint8, device="cuda"), + ), + scratch_cache=scratch_cache, + max_experts_per_batch=2, + ) + streams = [torch.cuda.Stream(), torch.cuda.Stream()] + + with torch.cuda.stream(streams[0]): + first_batch = [next(output) for _ in range(4)] + with torch.cuda.stream(streams[1]): + second_batch = [next(output) for _ in range(4)] + + assert len(first_batch) == len(second_batch) == 4 + assert len(scratch_cache) == 2 + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_refit_quantize_matches_receiver_quantize_mxfp8_weight(): """Sender prequantization and the receiver helper must agree bit-for-bit. From 97807c32c23e79e505a5a562af3ac1e008d45775 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Mon, 24 Aug 2026 14:35:50 -0700 Subject: [PATCH 62/76] fix(refit): synchronize batched outputs across streams Signed-off-by: seonjinn --- .../vllm/quantization/fp8_train_utils.py | 17 ++++++++- .../models/generation/test_mxfp8_prequant.py | 38 +++++++++++++++++++ 2 files changed, 53 insertions(+), 2 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py index 1c26e58312a..1acc080eafc 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py @@ -198,6 +198,9 @@ def flush_group( with torch.no_grad(): torch.stack(tensors, dim=0, out=stacked) + producer_stream = ( + torch.cuda.current_stream(first.device) if first.is_cuda else None + ) value, scale = quantize_fn(stacked.view(-1, stacked.shape[-1])) value = value.view_as(stacked) scale_columns = first.shape[-1] // MXFP8_BLOCK_SIZE @@ -208,8 +211,18 @@ def flush_group( ) scale = scale.view(len(chunk), *scale_shape) for index, (_expert_id, name, _tensor) in enumerate(chunk): - yield name, value[index] - yield name + "_scale_from_checkpoint", scale[index] + for output_name, output_tensor in ( + (name, value[index]), + (name + "_scale_from_checkpoint", scale[index]), + ): + if producer_stream is not None: + consumer_stream = torch.cuda.current_stream( + output_tensor.device + ) + if consumer_stream != producer_stream: + consumer_stream.wait_stream(producer_stream) + output_tensor.record_stream(consumer_stream) + yield output_name, output_tensor def flush_pending() -> Iterator[tuple[str, torch.Tensor]]: while pending: diff --git a/tests/unit/models/generation/test_mxfp8_prequant.py b/tests/unit/models/generation/test_mxfp8_prequant.py index e12c42e59ea..2d45c3cb126 100644 --- a/tests/unit/models/generation/test_mxfp8_prequant.py +++ b/tests/unit/models/generation/test_mxfp8_prequant.py @@ -348,6 +348,44 @@ def test_refit_quantize_matches_receiver_quantize_mxfp8_weight(): assert torch.equal(sent_scale, recv_scale) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_batched_expert_prequantization_waits_when_consumer_stream_changes(): + from nemo_rl.models.generation.vllm.quantization import fp8_train_utils + + def expert_name(expert_id): + return f"model.layers.0.mlp.experts.{expert_id}.gate_proj.weight" + + params = [ + (expert_name(i), torch.full((2, 64), i + 1, device="cuda")) for i in range(2) + ] + + def delayed_quantize(tensor): + value = torch.empty_like(tensor) + torch.cuda._sleep(5_000_000) + value.copy_(tensor) + scale = torch.ones((*tensor.shape[:-1], 2), dtype=torch.uint8, device="cuda") + return value, scale + + output = fp8_train_utils.iter_mxfp8_prequantized_params( + params, + {name for name, _tensor in params}, + quantize_fn=delayed_quantize, + max_experts_per_batch=2, + ) + producer_stream = torch.cuda.Stream() + consumer_stream = torch.cuda.Stream() + + with torch.cuda.stream(producer_stream): + next(output) + next(output) + with torch.cuda.stream(consumer_stream): + second_expert, _second_scale = next(output), next(output) + observed = second_expert[1].clone() + consumer_stream.synchronize() + + torch.testing.assert_close(observed, params[1][1]) + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") @pytest.mark.parametrize( "is_gated,intermediate_size,hidden_size", From 58316d1b860c7859200da6a26a14acdf5cc82e58 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Mon, 24 Aug 2026 14:45:45 -0700 Subject: [PATCH 63/76] fix(refit): synchronize fallback and pending tensors Signed-off-by: seonjinn --- .../vllm/quantization/fp8_train_utils.py | 85 +++++++++++++----- .../models/generation/test_mxfp8_prequant.py | 89 +++++++++++++++++++ 2 files changed, 150 insertions(+), 24 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py index 1acc080eafc..8753ec60f92 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py @@ -136,17 +136,49 @@ def iter_mxfp8_prequantized_params( if scratch_cache is None: scratch_cache = {} - pending: dict[tuple[str, str], list[tuple[int, str, torch.Tensor]]] = {} + pending: dict[ + tuple[str, str], + list[tuple[int, str, torch.Tensor, torch.cuda.Stream | None]], + ] = {} current_prefix: str | None = None - def quantize_one(name: str, tensor: torch.Tensor) -> list[tuple[str, torch.Tensor]]: + def yield_on_current_stream( + entries: Iterable[tuple[str, torch.Tensor]], + producer_stream: torch.cuda.Stream | None, + ) -> Iterator[tuple[str, torch.Tensor]]: + for output_name, output_tensor in entries: + if producer_stream is not None: + consumer_stream = torch.cuda.current_stream(output_tensor.device) + if consumer_stream != producer_stream: + consumer_stream.wait_stream(producer_stream) + output_tensor.record_stream(consumer_stream) + yield output_name, output_tensor + + def quantize_one( + name: str, + tensor: torch.Tensor, + source_stream: torch.cuda.Stream | None = None, + ) -> Iterator[tuple[str, torch.Tensor]]: if tensor.dtype == torch.float8_e4m3fn: raise ValueError( "MXFP8 prequantization requires BF16 trainer-exported weights; " f"{name} is already stored as E4M3." ) + producer_stream = ( + torch.cuda.current_stream(tensor.device) if tensor.is_cuda else None + ) + if ( + producer_stream is not None + and source_stream is not None + and source_stream != producer_stream + ): + producer_stream.wait_stream(source_stream) + tensor.record_stream(producer_stream) value, scale = quantize_fn(tensor) - return [(name, value), (name + "_scale_from_checkpoint", scale)] + yield from yield_on_current_stream( + ((name, value), (name + "_scale_from_checkpoint", scale)), + producer_stream, + ) def flush_group( group_key: tuple[str, str], @@ -156,7 +188,7 @@ def flush_group( while group: chunk = group[:max_experts_per_batch] del group[:max_experts_per_batch] - tensors = [tensor for _expert_id, _name, tensor in chunk] + tensors = [tensor for _expert_id, _name, tensor, _stream in chunk] batchable = len(chunk) > 1 and len({item[0] for item in chunk}) == len( chunk ) @@ -170,8 +202,8 @@ def flush_group( for tensor in tensors ) if not batchable: - for _expert_id, name, tensor in chunk: - yield from quantize_one(name, tensor) + for _expert_id, name, tensor, source_stream in chunk: + yield from quantize_one(name, tensor, source_stream) continue first = tensors[0] @@ -180,10 +212,18 @@ def flush_group( "MXFP8 prequantization requires BF16 trainer-exported weights." ) required_numel = len(chunk) * first.numel() + stack_stream = ( + torch.cuda.current_stream(first.device) if first.is_cuda else None + ) + if stack_stream is not None: + for tensor, (_expert_id, _name, _tensor, source_stream) in zip( + tensors, chunk + ): + if source_stream is not None and source_stream != stack_stream: + stack_stream.wait_stream(source_stream) + tensor.record_stream(stack_stream) stream_id = ( - int(torch.cuda.current_stream(first.device).cuda_stream) - if first.is_cuda - else None + int(stack_stream.cuda_stream) if stack_stream is not None else None ) cache_key = (first.device, first.dtype, stream_id) scratch = scratch_cache.get(cache_key) @@ -198,9 +238,7 @@ def flush_group( with torch.no_grad(): torch.stack(tensors, dim=0, out=stacked) - producer_stream = ( - torch.cuda.current_stream(first.device) if first.is_cuda else None - ) + producer_stream = stack_stream value, scale = quantize_fn(stacked.view(-1, stacked.shape[-1])) value = value.view_as(stacked) scale_columns = first.shape[-1] // MXFP8_BLOCK_SIZE @@ -210,19 +248,15 @@ def flush_group( else (*first.shape[:-1], scale_columns) ) scale = scale.view(len(chunk), *scale_shape) - for index, (_expert_id, name, _tensor) in enumerate(chunk): - for output_name, output_tensor in ( + entries = ( + entry + for index, (_expert_id, name, _tensor, _stream) in enumerate(chunk) + for entry in ( (name, value[index]), (name + "_scale_from_checkpoint", scale[index]), - ): - if producer_stream is not None: - consumer_stream = torch.cuda.current_stream( - output_tensor.device - ) - if consumer_stream != producer_stream: - consumer_stream.wait_stream(producer_stream) - output_tensor.record_stream(consumer_stream) - yield output_name, output_tensor + ) + ) + yield from yield_on_current_stream(entries, producer_stream) def flush_pending() -> Iterator[tuple[str, torch.Tensor]]: while pending: @@ -247,7 +281,10 @@ def flush_pending() -> Iterator[tuple[str, torch.Tensor]]: current_prefix = prefix group_key = (prefix, projection) group = pending.setdefault(group_key, []) - group.append((int(match.group("expert_id")), name, tensor)) + source_stream = ( + torch.cuda.current_stream(tensor.device) if tensor.is_cuda else None + ) + group.append((int(match.group("expert_id")), name, tensor, source_stream)) if len(group) == max_experts_per_batch: yield from flush_group(group_key) diff --git a/tests/unit/models/generation/test_mxfp8_prequant.py b/tests/unit/models/generation/test_mxfp8_prequant.py index 2d45c3cb126..3d53305153f 100644 --- a/tests/unit/models/generation/test_mxfp8_prequant.py +++ b/tests/unit/models/generation/test_mxfp8_prequant.py @@ -386,6 +386,95 @@ def delayed_quantize(tensor): torch.testing.assert_close(observed, params[1][1]) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_prequantization_fallback_waits_when_consumer_stream_changes(): + from nemo_rl.models.generation.vllm.quantization import fp8_train_utils + + name = "model.layers.0.self_attn.q_proj.weight" + tensor = torch.full((2, 64), 7, device="cuda") + + def delayed_quantize(input_tensor): + value = torch.empty_like(input_tensor) + torch.cuda._sleep(5_000_000) + value.copy_(input_tensor) + scale = torch.full( + (*input_tensor.shape[:-1], 2), 3, dtype=torch.uint8, device="cuda" + ) + return value, scale + + output = fp8_train_utils.iter_mxfp8_prequantized_params( + [(name, tensor)], + {name}, + quantize_fn=delayed_quantize, + ) + producer_stream = torch.cuda.Stream() + consumer_stream = torch.cuda.Stream() + + with torch.cuda.stream(producer_stream): + next(output) + with torch.cuda.stream(consumer_stream): + scale_name, scale = next(output) + observed = scale.clone() + consumer_stream.synchronize() + + assert scale_name == name + "_scale_from_checkpoint" + assert torch.all(observed == 3) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize("second_up_shape", [(2, 64), (3, 64)]) +def test_batched_expert_prequantization_waits_for_pending_input_stream( + second_up_shape, +): + from nemo_rl.models.generation.vllm.quantization import fp8_train_utils + + def expert_name(expert_id, projection): + return f"model.layers.0.mlp.experts.{expert_id}.{projection}_proj.weight" + + def params(): + yield expert_name(0, "gate"), torch.ones(2, 64, device="cuda") + delayed_up = torch.empty(2, 64, device="cuda") + torch.cuda._sleep(5_000_000) + delayed_up.fill_(7) + yield expert_name(0, "up"), delayed_up + yield expert_name(0, "down"), torch.ones(4, 32, device="cuda") + yield expert_name(1, "gate"), torch.full((2, 64), 2, device="cuda") + yield expert_name(1, "up"), torch.full(second_up_shape, 3, device="cuda") + yield expert_name(1, "down"), torch.full((4, 32), 4, device="cuda") + + selected_names = { + expert_name(expert_id, projection) + for expert_id in range(2) + for projection in ("gate", "up", "down") + } + output = fp8_train_utils.iter_mxfp8_prequantized_params( + params(), + selected_names, + quantize_fn=lambda input_tensor: ( + input_tensor.clone(), + torch.ones( + (*input_tensor.shape[:-1], input_tensor.shape[-1] // 32), + dtype=torch.uint8, + device="cuda", + ), + ), + max_experts_per_batch=2, + ) + producer_stream = torch.cuda.Stream() + consumer_stream = torch.cuda.Stream() + + with torch.cuda.stream(producer_stream): + gate_entries = [next(output) for _ in range(4)] + with torch.cuda.stream(consumer_stream): + up_name, up_tensor = next(output) + observed = up_tensor.clone() + consumer_stream.synchronize() + + assert len(gate_entries) == 4 + assert up_name == expert_name(0, "up") + torch.testing.assert_close(observed, torch.full_like(observed, 7)) + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") @pytest.mark.parametrize( "is_gated,intermediate_size,hidden_size", From cb025c6aa0ef749819d579bae5a40b489c4bb97f Mon Sep 17 00:00:00 2001 From: seonjinn Date: Mon, 24 Aug 2026 15:02:56 -0700 Subject: [PATCH 64/76] test(refit): use BF16 grad-enabled expert fixtures Signed-off-by: seonjinn --- tests/unit/models/generation/test_mxfp8_prequant.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/tests/unit/models/generation/test_mxfp8_prequant.py b/tests/unit/models/generation/test_mxfp8_prequant.py index 3d53305153f..ab2b32e29a2 100644 --- a/tests/unit/models/generation/test_mxfp8_prequant.py +++ b/tests/unit/models/generation/test_mxfp8_prequant.py @@ -153,7 +153,12 @@ def expert_name(expert_id, projection): else: shape = (2, 64) fill_value = expert_id + (1 if projection == "gate" else 3) - tensor = torch.full(shape, fill_value, requires_grad=True) + tensor = torch.full( + shape, + fill_value, + dtype=torch.bfloat16, + requires_grad=True, + ) params.append((name, tensor)) expected[name] = tensor From 224030a503192949efb3b0c2735c8487fb73f743 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Mon, 24 Aug 2026 21:20:30 -0700 Subject: [PATCH 65/76] test(refit): cover split MXFP8 weight and scale batches Signed-off-by: seonjinn --- .../generation/test_vllm_fp8_quantization.py | 37 ++++++++++++++----- 1 file changed, 27 insertions(+), 10 deletions(-) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 31735b52bde..b58ae6a75c0 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -1119,22 +1119,39 @@ def test_load_weights_rejects_unnegotiated_mxfp8_payload(fp8_module, monkeypatch ) -def test_load_weights_rejects_prequantized_mxfp8_without_scale(fp8_module, monkeypatch): +def test_load_weights_accepts_prequantized_mxfp8_split_across_batches( + fp8_module, monkeypatch +): + from nemo_rl.models.generation.vllm import vllm_backend + fp8 = fp8_module fp8.global_fp8_config = types.SimpleNamespace( is_mx=True, refit_prequantize=True, ) - monkeypatch.setattr(fp8, "_is_fp8_weight", lambda _name, _model: True) + weight = torch.ones(2, 2, dtype=torch.float8_e4m3fn) + scale = torch.ones(2, 1, dtype=torch.uint8) + loaded = [] + monkeypatch.setattr( + fp8, "_is_fp8_weight", lambda name, _model: name.endswith(".weight") + ) + monkeypatch.setattr( + vllm_backend, + "load_weights_maybe_cached", + lambda model, weights, *, cache_loader_routes: loaded.extend(weights), + ) + model_runner = types.SimpleNamespace( + model=object(), + vllm_config=types.SimpleNamespace(additional_config={}), + ) - with pytest.raises(ValueError, match="missing.*scale_from_checkpoint"): - fp8.load_weights( - [("model.weight", torch.ones(2, 2, dtype=torch.float8_e4m3fn))], - types.SimpleNamespace( - model=object(), - vllm_config=types.SimpleNamespace(additional_config={}), - ), - ) + fp8.load_weights([("model.weight", weight)], model_runner) + fp8.load_weights([("model.weight_scale_from_checkpoint", scale)], model_runner) + + assert loaded == [ + ["model.weight", weight], + ("model.weight_scale_from_checkpoint", scale), + ] def test_load_weights_preserves_non_mx_blockwise_fp8_payload(fp8_module, monkeypatch): From d234588bc3d4912ab5fea87185764aceb93bc19d Mon Sep 17 00:00:00 2001 From: seonjinn Date: Mon, 24 Aug 2026 21:25:18 -0700 Subject: [PATCH 66/76] fix(refit): allow split MXFP8 weight scale batches Signed-off-by: seonjinn --- nemo_rl/models/generation/vllm/quantization/fp8.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index c96a7e99ba1..5d740c5aed1 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -678,7 +678,10 @@ def get_quantized_weight_iterator( f"Prequantized MXFP8 weight {k!r} is missing {scale_name!r}." ) # Prequantized MXFP8 sends the matching *_scale_from_checkpoint - # entry separately. Non-MXFP8 blockwise FP8 sends *_scale_inv. + # entry separately, and IPC buffer boundaries may place that scale + # in the next batch. The IPC manifest validates the complete set + # before post-load processing. Non-MXFP8 blockwise FP8 sends + # *_scale_inv. yield k, v continue is_mx = global_fp8_config.is_mx From be3de0200825cd88111b9151505ba2623c03ce7b Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sat, 5 Sep 2026 01:25:05 -0700 Subject: [PATCH 67/76] test(refit): reject MXFP8 parameter prequantization Signed-off-by: seonjinn --- tests/unit/models/policy/test_megatron_worker.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index c3f9aaa7b07..dfeb4f69227 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -823,7 +823,8 @@ def iter_batched(params, selected_names): assert calls[0][1] == {name} -def test_enable_refit_prequantize_rejects_blockwise_fp8_storage(): +@pytest.mark.parametrize("fp8_recipe", ["blockwise", "mxfp8"]) +def test_enable_refit_prequantize_rejects_fp8_param_storage(fp8_recipe): from nemo_rl.models.policy.workers.megatron_policy_worker import ( MegatronPolicyWorkerImpl, ) @@ -831,7 +832,7 @@ def test_enable_refit_prequantize_rejects_blockwise_fp8_storage(): worker = object.__new__(MegatronPolicyWorkerImpl) worker.fp8_cfg = { "fp8_param": True, - "fp8_recipe": "blockwise", + "fp8_recipe": fp8_recipe, } with pytest.raises(ValueError, match="BF16 trainer-exported weights"): From 2e5507c8c56ca060ed7a41397ddf056ac563d08a Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sat, 5 Sep 2026 01:48:03 -0700 Subject: [PATCH 68/76] test(refit): cover batched prequantization contracts Signed-off-by: seonjinn --- .../models/generation/test_mxfp8_prequant.py | 28 +++++-------------- .../generation/test_vllm_fp8_quantization.py | 23 +++++++++++++++ .../models/policy/test_megatron_worker.py | 4 ++- 3 files changed, 33 insertions(+), 22 deletions(-) diff --git a/tests/unit/models/generation/test_mxfp8_prequant.py b/tests/unit/models/generation/test_mxfp8_prequant.py index ab2b32e29a2..2d72cc5bf1f 100644 --- a/tests/unit/models/generation/test_mxfp8_prequant.py +++ b/tests/unit/models/generation/test_mxfp8_prequant.py @@ -179,11 +179,7 @@ def expert_name(expert_id, projection): torch.testing.assert_close(output[name], tensor) scale_name = name + "_scale_from_checkpoint" scale_columns = tensor.shape[-1] // MXFP8_BLOCK_SIZE - expected_scale_shape = ( - tensor.shape[:-1] - if scale_columns == 1 - else (*tensor.shape[:-1], scale_columns) - ) + expected_scale_shape = (*tensor.shape[:-1], scale_columns) assert output[scale_name].shape == expected_scale_shape assert torch.all(output[scale_name] == 1) @@ -202,7 +198,7 @@ def expert_name(expert_id, projection): assert next(iter(scratch_cache.values())).data_ptr() == first_scratch_ptr -def test_batched_expert_prequantization_bounds_batch_and_has_stable_order(): +def test_batched_expert_prequantization_bounds_batch_and_preserves_source_order(): from nemo_rl.models.generation.vllm.quantization import fp8_train_utils calls = [] @@ -233,21 +229,11 @@ def expert_name(expert_id, projection): ) ) - expected_names = [] - for expert_ids, projection in ( - ((0, 1), "gate"), - ((0, 1), "up"), - ((0, 1), "down"), - ((2, 3), "gate"), - ((2, 3), "up"), - ((2, 3), "down"), - ((4,), "gate"), - ((4,), "up"), - ((4,), "down"), - ): - for expert_id in expert_ids: - name = expert_name(expert_id, projection) - expected_names.extend((name, name + "_scale_from_checkpoint")) + expected_names = [ + output_name + for name, _tensor in params + for output_name in (name, name + "_scale_from_checkpoint") + ] assert [name for name, _tensor in output] == expected_names assert calls == [ diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index b58ae6a75c0..c3c6482f231 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -1131,6 +1131,9 @@ def test_load_weights_accepts_prequantized_mxfp8_split_across_batches( ) weight = torch.ones(2, 2, dtype=torch.float8_e4m3fn) scale = torch.ones(2, 1, dtype=torch.uint8) + fp8.set_refit_manifest_names( + {"model.weight", "model.weight_scale_from_checkpoint"} + ) loaded = [] monkeypatch.setattr( fp8, "_is_fp8_weight", lambda name, _model: name.endswith(".weight") @@ -1154,6 +1157,26 @@ def test_load_weights_accepts_prequantized_mxfp8_split_across_batches( ] +def test_load_weights_rejects_prequantized_mxfp8_without_scale_or_manifest( + fp8_module, monkeypatch +): + fp8 = fp8_module + fp8.global_fp8_config = types.SimpleNamespace( + is_mx=True, + refit_prequantize=True, + ) + monkeypatch.setattr(fp8, "_is_fp8_weight", lambda _name, _model: True) + + with pytest.raises(ValueError, match="missing.*scale_from_checkpoint"): + fp8.load_weights( + [("model.weight", torch.ones(2, 2, dtype=torch.float8_e4m3fn))], + types.SimpleNamespace( + model=object(), + vllm_config=types.SimpleNamespace(additional_config={}), + ), + ) + + def test_load_weights_preserves_non_mx_blockwise_fp8_payload(fp8_module, monkeypatch): from nemo_rl.models.generation.vllm import vllm_backend diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index dfeb4f69227..1ba828a8e85 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -824,7 +824,9 @@ def iter_batched(params, selected_names): @pytest.mark.parametrize("fp8_recipe", ["blockwise", "mxfp8"]) -def test_enable_refit_prequantize_rejects_fp8_param_storage(fp8_recipe): +def test_enable_refit_prequantize_rejects_fp8_param_storage( + fp8_recipe: str, +) -> None: from nemo_rl.models.policy.workers.megatron_policy_worker import ( MegatronPolicyWorkerImpl, ) From 4dff3ccdcf781e5f724cc05a20e9d57a003c85b3 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sat, 5 Sep 2026 02:13:45 -0700 Subject: [PATCH 69/76] fix(refit): preserve batched prequant wire order Signed-off-by: seonjinn --- .../generation/vllm/quantization/fp8.py | 2 + .../vllm/quantization/fp8_train_utils.py | 196 ++++++++++-------- .../policy/workers/megatron_policy_worker.py | 4 +- .../models/policy/test_megatron_worker.py | 42 ++++ 4 files changed, 155 insertions(+), 89 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index 5d740c5aed1..d01725c9973 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -614,6 +614,8 @@ def get_quantized_weight_iterator( ) -> Iterator[tuple[str, torch.Tensor]]: """Convert trainer weights to the checkpoint tensors expected by vLLM.""" model = model_runner.model + weights = list(weights) + weight_names = {name for name, _tensor in weights} for k, v in weights: grouped_weight_name = _grouped_expert_weight_name_from_scale(k) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py index 8753ec60f92..2b52aa04333 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py @@ -136,10 +136,10 @@ def iter_mxfp8_prequantized_params( if scratch_cache is None: scratch_cache = {} - pending: dict[ - tuple[str, str], - list[tuple[int, str, torch.Tensor, torch.cuda.Stream | None]], - ] = {} + pending: list[ + tuple[int, str, str, torch.Tensor, torch.cuda.Stream | None] + ] = [] + pending_expert_ids: set[int] = set() current_prefix: str | None = None def yield_on_current_stream( @@ -154,11 +154,11 @@ def yield_on_current_stream( output_tensor.record_stream(consumer_stream) yield output_name, output_tensor - def quantize_one( + def quantize_one_result( name: str, tensor: torch.Tensor, source_stream: torch.cuda.Stream | None = None, - ) -> Iterator[tuple[str, torch.Tensor]]: + ) -> tuple[torch.Tensor, torch.Tensor, torch.cuda.Stream | None]: if tensor.dtype == torch.float8_e4m3fn: raise ValueError( "MXFP8 prequantization requires BF16 trainer-exported weights; " @@ -175,92 +175,111 @@ def quantize_one( producer_stream.wait_stream(source_stream) tensor.record_stream(producer_stream) value, scale = quantize_fn(tensor) + return value, scale, producer_stream + + def quantize_one( + name: str, + tensor: torch.Tensor, + source_stream: torch.cuda.Stream | None = None, + ) -> Iterator[tuple[str, torch.Tensor]]: + value, scale, producer_stream = quantize_one_result( + name, tensor, source_stream + ) yield from yield_on_current_stream( ((name, value), (name + "_scale_from_checkpoint", scale)), producer_stream, ) - def flush_group( - group_key: tuple[str, str], - ) -> Iterator[tuple[str, torch.Tensor]]: - group = pending.pop(group_key) - group.sort(key=lambda item: item[0]) - while group: - chunk = group[:max_experts_per_batch] - del group[:max_experts_per_batch] - tensors = [tensor for _expert_id, _name, tensor, _stream in chunk] - batchable = len(chunk) > 1 and len({item[0] for item in chunk}) == len( - chunk - ) - if batchable: + def flush_pending() -> Iterator[tuple[str, torch.Tensor]]: + if not pending: + return + + results: dict[ + int, tuple[torch.Tensor, torch.Tensor, torch.cuda.Stream | None] + ] = {} + projection_groups: dict[str, list[int]] = {} + for index, (_expert_id, projection, _name, _tensor, _stream) in enumerate( + pending + ): + projection_groups.setdefault(projection, []).append(index) + + for indices in projection_groups.values(): + for chunk_start in range(0, len(indices), max_experts_per_batch): + chunk_indices = indices[ + chunk_start : chunk_start + max_experts_per_batch + ] + chunk = [pending[index] for index in chunk_indices] + tensors = [tensor for _id, _proj, _name, tensor, _stream in chunk] + batchable = len(chunk) > 1 and len( + {item[0] for item in chunk} + ) == len(chunk) + if batchable: + first = tensors[0] + batchable = all( + tensor.shape == first.shape + and tensor.dtype == first.dtype + and tensor.device == first.device + and tensor.layout is torch.strided + for tensor in tensors + ) + if not batchable: + for index, (_id, _proj, name, tensor, source_stream) in zip( + chunk_indices, chunk + ): + results[index] = quantize_one_result( + name, tensor, source_stream + ) + continue + first = tensors[0] - batchable = all( - tensor.shape == first.shape - and tensor.dtype == first.dtype - and tensor.device == first.device - and tensor.layout is torch.strided - for tensor in tensors - ) - if not batchable: - for _expert_id, name, tensor, source_stream in chunk: - yield from quantize_one(name, tensor, source_stream) - continue - - first = tensors[0] - if first.dtype == torch.float8_e4m3fn: - raise ValueError( - "MXFP8 prequantization requires BF16 trainer-exported weights." + if first.dtype == torch.float8_e4m3fn: + raise ValueError( + "MXFP8 prequantization requires BF16 trainer-exported weights." + ) + required_numel = len(chunk) * first.numel() + stack_stream = ( + torch.cuda.current_stream(first.device) if first.is_cuda else None ) - required_numel = len(chunk) * first.numel() - stack_stream = ( - torch.cuda.current_stream(first.device) if first.is_cuda else None - ) - if stack_stream is not None: - for tensor, (_expert_id, _name, _tensor, source_stream) in zip( - tensors, chunk - ): - if source_stream is not None and source_stream != stack_stream: - stack_stream.wait_stream(source_stream) - tensor.record_stream(stack_stream) - stream_id = ( - int(stack_stream.cuda_stream) if stack_stream is not None else None - ) - cache_key = (first.device, first.dtype, stream_id) - scratch = scratch_cache.get(cache_key) - if scratch is None or scratch.numel() < required_numel: - scratch = torch.empty( - required_numel, - dtype=first.dtype, - device=first.device, + if stack_stream is not None: + for tensor, (_id, _proj, _name, _tensor, source_stream) in zip( + tensors, chunk + ): + if source_stream is not None and source_stream != stack_stream: + stack_stream.wait_stream(source_stream) + tensor.record_stream(stack_stream) + stream_id = ( + int(stack_stream.cuda_stream) if stack_stream is not None else None ) - scratch_cache[cache_key] = scratch - stacked = scratch[:required_numel].view(len(chunk), *first.shape) - with torch.no_grad(): - torch.stack(tensors, dim=0, out=stacked) - - producer_stream = stack_stream - value, scale = quantize_fn(stacked.view(-1, stacked.shape[-1])) - value = value.view_as(stacked) - scale_columns = first.shape[-1] // MXFP8_BLOCK_SIZE - scale_shape = ( - first.shape[:-1] - if scale_columns == 1 - else (*first.shape[:-1], scale_columns) + cache_key = (first.device, first.dtype, stream_id) + scratch = scratch_cache.get(cache_key) + if scratch is None or scratch.numel() < required_numel: + scratch = torch.empty( + required_numel, + dtype=first.dtype, + device=first.device, + ) + scratch_cache[cache_key] = scratch + stacked = scratch[:required_numel].view(len(chunk), *first.shape) + with torch.no_grad(): + torch.stack(tensors, dim=0, out=stacked) + + value, scale = quantize_fn(stacked.view(-1, stacked.shape[-1])) + value = value.view_as(stacked) + scale_columns = first.shape[-1] // MXFP8_BLOCK_SIZE + scale_shape = (*first.shape[:-1], scale_columns) + scale = scale.view(len(chunk), *scale_shape) + for offset, index in enumerate(chunk_indices): + results[index] = value[offset], scale[offset], stack_stream + + for index, (_id, _proj, name, _tensor, _stream) in enumerate(pending): + value, scale, producer_stream = results[index] + yield from yield_on_current_stream( + ((name, value), (name + "_scale_from_checkpoint", scale)), + producer_stream, ) - scale = scale.view(len(chunk), *scale_shape) - entries = ( - entry - for index, (_expert_id, name, _tensor, _stream) in enumerate(chunk) - for entry in ( - (name, value[index]), - (name + "_scale_from_checkpoint", scale[index]), - ) - ) - yield from yield_on_current_stream(entries, producer_stream) - def flush_pending() -> Iterator[tuple[str, torch.Tensor]]: - while pending: - yield from flush_group(next(iter(pending))) + pending.clear() + pending_expert_ids.clear() for name, tensor in params: match = _EXPERT_WEIGHT_PATTERN.match(name) if name in selected_names else None @@ -279,14 +298,17 @@ def flush_pending() -> Iterator[tuple[str, torch.Tensor]]: if current_prefix is not None and prefix != current_prefix: yield from flush_pending() current_prefix = prefix - group_key = (prefix, projection) - group = pending.setdefault(group_key, []) + expert_id = int(match.group("expert_id")) + if ( + expert_id not in pending_expert_ids + and len(pending_expert_ids) == max_experts_per_batch + ): + yield from flush_pending() source_stream = ( torch.cuda.current_stream(tensor.device) if tensor.is_cuda else None ) - group.append((int(match.group("expert_id")), name, tensor, source_stream)) - if len(group) == max_experts_per_batch: - yield from flush_group(group_key) + pending.append((expert_id, projection, name, tensor, source_stream)) + pending_expert_ids.add(expert_id) if pending: yield from flush_pending() diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index f6e21ae5ffb..b54d2440908 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -2351,10 +2351,10 @@ def enable_refit_prequantize(self, param_names: list[str]) -> dict[str, Any]: Updated refit metadata: the listed params become float8_e4m3fn and each gains a *_scale_from_checkpoint uint8 entry. """ - if self._is_fp8_export(): + if self.fp8_cfg is not None and self.fp8_cfg.get("fp8_param", False): raise ValueError( "vllm_cfg.refit_prequantize requires BF16 trainer-exported weights; " - "Megatron blockwise FP8 parameter storage uses a different scale layout." + "Megatron FP8 parameter storage uses a different scale layout." ) if self._refit_param_info_hf is None: raise RuntimeError( diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index 1ba828a8e85..22d037f28cd 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -823,6 +823,48 @@ def iter_batched(params, selected_names): assert calls[0][1] == {name} +def test_iter_params_preserves_bridge_expert_wire_order(monkeypatch): + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + + names = [ + f"model.layers.0.mlp.experts.{expert_id}.gate_proj.weight" + for expert_id in range(2) + ] + weights = [ + torch.full((2, 32), expert_id + 1, dtype=torch.bfloat16) + for expert_id in range(2) + ] + stack_calls = [] + original_stack = torch.stack + + def record_stack(tensors, *args, **kwargs): + stack_calls.append([tensor.clone() for tensor in tensors]) + return original_stack(tensors, *args, **kwargs) + + monkeypatch.setattr(torch, "stack", record_stack) + worker = object.__new__(MegatronPolicyWorkerImpl) + worker._refit_prequant_names = set(names) + worker.model = object() + worker.draft_model = None + worker.refit_conversion_tasks = [] + worker.cfg = {"megatron_cfg": {"enabled": True}} + worker.megatron_bridge = SimpleNamespace( + export_hf_weights=lambda *_args, **_kwargs: iter(zip(names, weights)) + ) + + output = list(worker._iter_params_with_optional_kv_scales()) + + assert [name for name, _tensor in output] == [ + entry_name + for name in names + for entry_name in (name, name + "_scale_from_checkpoint") + ] + assert len(stack_calls) == 1 + assert len(stack_calls[0]) == 2 + + @pytest.mark.parametrize("fp8_recipe", ["blockwise", "mxfp8"]) def test_enable_refit_prequantize_rejects_fp8_param_storage( fp8_recipe: str, From 05889896f234cb60b2ae3125aa6ecc40d54db9d4 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sat, 5 Sep 2026 02:18:04 -0700 Subject: [PATCH 70/76] style: format refit prequantization changes Signed-off-by: seonjinn --- .../vllm/quantization/fp8_train_utils.py | 14 +++++--------- .../generation/test_vllm_fp8_quantization.py | 4 +--- 2 files changed, 6 insertions(+), 12 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py index 2b52aa04333..d063026de5e 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py @@ -136,9 +136,7 @@ def iter_mxfp8_prequantized_params( if scratch_cache is None: scratch_cache = {} - pending: list[ - tuple[int, str, str, torch.Tensor, torch.cuda.Stream | None] - ] = [] + pending: list[tuple[int, str, str, torch.Tensor, torch.cuda.Stream | None]] = [] pending_expert_ids: set[int] = set() current_prefix: str | None = None @@ -182,9 +180,7 @@ def quantize_one( tensor: torch.Tensor, source_stream: torch.cuda.Stream | None = None, ) -> Iterator[tuple[str, torch.Tensor]]: - value, scale, producer_stream = quantize_one_result( - name, tensor, source_stream - ) + value, scale, producer_stream = quantize_one_result(name, tensor, source_stream) yield from yield_on_current_stream( ((name, value), (name + "_scale_from_checkpoint", scale)), producer_stream, @@ -210,9 +206,9 @@ def flush_pending() -> Iterator[tuple[str, torch.Tensor]]: ] chunk = [pending[index] for index in chunk_indices] tensors = [tensor for _id, _proj, _name, tensor, _stream in chunk] - batchable = len(chunk) > 1 and len( - {item[0] for item in chunk} - ) == len(chunk) + batchable = len(chunk) > 1 and len({item[0] for item in chunk}) == len( + chunk + ) if batchable: first = tensors[0] batchable = all( diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index c3c6482f231..6766654ce61 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -1131,9 +1131,7 @@ def test_load_weights_accepts_prequantized_mxfp8_split_across_batches( ) weight = torch.ones(2, 2, dtype=torch.float8_e4m3fn) scale = torch.ones(2, 1, dtype=torch.uint8) - fp8.set_refit_manifest_names( - {"model.weight", "model.weight_scale_from_checkpoint"} - ) + fp8.set_refit_manifest_names({"model.weight", "model.weight_scale_from_checkpoint"}) loaded = [] monkeypatch.setattr( fp8, "_is_fp8_weight", lambda name, _model: name.endswith(".weight") From c34e554f77fe5d2711ffe22bc2b641deffac8a37 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sat, 5 Sep 2026 02:23:13 -0700 Subject: [PATCH 71/76] test(refit): follow preserved expert wire order Signed-off-by: seonjinn --- tests/unit/models/generation/test_mxfp8_prequant.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/unit/models/generation/test_mxfp8_prequant.py b/tests/unit/models/generation/test_mxfp8_prequant.py index 2d72cc5bf1f..7725f8822ba 100644 --- a/tests/unit/models/generation/test_mxfp8_prequant.py +++ b/tests/unit/models/generation/test_mxfp8_prequant.py @@ -455,13 +455,13 @@ def params(): consumer_stream = torch.cuda.Stream() with torch.cuda.stream(producer_stream): - gate_entries = [next(output) for _ in range(4)] + gate_entries = [next(output) for _ in range(2)] with torch.cuda.stream(consumer_stream): up_name, up_tensor = next(output) observed = up_tensor.clone() consumer_stream.synchronize() - assert len(gate_entries) == 4 + assert len(gate_entries) == 2 assert up_name == expert_name(0, "up") torch.testing.assert_close(observed, torch.full_like(observed, 7)) From b786b3996043533587d9bf5b653f45f7579a2851 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sat, 5 Sep 2026 02:26:13 -0700 Subject: [PATCH 72/76] fix(refit): honor disabled FP8 parameter config Signed-off-by: seonjinn --- .../policy/workers/megatron_policy_worker.py | 6 ++++- .../models/policy/test_megatron_worker.py | 25 +++++++++++++++++++ 2 files changed, 30 insertions(+), 1 deletion(-) diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index b54d2440908..7c6c744ed69 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -2351,7 +2351,11 @@ def enable_refit_prequantize(self, param_names: list[str]) -> dict[str, Any]: Updated refit metadata: the listed params become float8_e4m3fn and each gains a *_scale_from_checkpoint uint8 entry. """ - if self.fp8_cfg is not None and self.fp8_cfg.get("fp8_param", False): + if ( + self.fp8_cfg is not None + and self.fp8_cfg.get("enabled", False) + and self.fp8_cfg.get("fp8_param", False) + ): raise ValueError( "vllm_cfg.refit_prequantize requires BF16 trainer-exported weights; " "Megatron FP8 parameter storage uses a different scale layout." diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index 22d037f28cd..21dbc5815e8 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -875,6 +875,7 @@ def test_enable_refit_prequantize_rejects_fp8_param_storage( worker = object.__new__(MegatronPolicyWorkerImpl) worker.fp8_cfg = { + "enabled": True, "fp8_param": True, "fp8_recipe": fp8_recipe, } @@ -883,6 +884,30 @@ def test_enable_refit_prequantize_rejects_fp8_param_storage( worker.enable_refit_prequantize(["model.weight"]) +def test_enable_refit_prequantize_allows_disabled_fp8_param_storage() -> None: + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + + worker = object.__new__(MegatronPolicyWorkerImpl) + worker.fp8_cfg = { + "enabled": False, + "fp8_param": True, + "fp8_recipe": "mxfp8", + } + worker._refit_param_info_hf = { + "model.weight": (torch.Size([4, 64]), torch.bfloat16), + } + + info = worker.enable_refit_prequantize(["model.weight"]) + + assert info["model.weight"] == (torch.Size([4, 64]), torch.float8_e4m3fn) + assert info["model.weight_scale_from_checkpoint"] == ( + torch.Size([4, 2]), + torch.uint8, + ) + + def test_enable_refit_prequantize_requires_prepare_refit_info(): from nemo_rl.models.policy.workers.megatron_policy_worker import ( MegatronPolicyWorkerImpl, From 238f2415458bede0b1a9d53be1f2c660f061fa89 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sun, 6 Sep 2026 22:47:03 -0700 Subject: [PATCH 73/76] fix(refit): preserve MXFP8 kernel scale storage Signed-off-by: seonjinn --- .../generation/vllm/quantization/fp8.py | 15 +++++++-- .../generation/test_vllm_fp8_quantization.py | 33 +++++++++++++++++-- 2 files changed, 43 insertions(+), 5 deletions(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8.py b/nemo_rl/models/generation/vllm/quantization/fp8.py index c96a7e99ba1..5c408b5953e 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8.py @@ -1726,8 +1726,19 @@ def process_weights_after_loading_mxfp8_moe(self, layer: RoutedExperts) -> None: ) else: assert self.moe_quant_config is not None - assert self.moe_quant_config.w1_scale is runtime_w13_scale - assert self.moe_quant_config.w2_scale is runtime_w2_scale + for kernel_scale, runtime_scale, scale_name in ( + (self.moe_quant_config.w1_scale, runtime_w13_scale, "w13"), + (self.moe_quant_config.w2_scale, runtime_w2_scale, "w2"), + ): + if kernel_scale is runtime_scale: + continue + if kernel_scale.shape != runtime_scale.shape: + raise RuntimeError( + f"MXFP8 MoE {scale_name} runtime scale shape changed from " + f"{tuple(kernel_scale.shape)} to {tuple(runtime_scale.shape)}" + ) + # Keep the storage already captured by the kernel and CUDA Graph. + kernel_scale.copy_(runtime_scale) def apply_monolithic_mxfp8_moe( diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 31735b52bde..9d553f86cef 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -777,7 +777,10 @@ def test_process_mxfp8_moe_refit_rejects_non_flashinfer_backend(fp8_module): fp8_module.process_weights_after_loading_mxfp8_moe(quant_method, object()) -def test_process_mxfp8_moe_initializes_kernel_once(fp8_module, monkeypatch): +@pytest.mark.parametrize("replace_runtime_scales", [False, True]) +def test_process_mxfp8_moe_initializes_kernel_once( + fp8_module, monkeypatch, replace_runtime_scales +): from vllm.model_executor.layers.fused_moe.oracle.fp8 import Fp8MoeBackend fp8 = fp8_module @@ -857,6 +860,13 @@ def make_kernel(**kwargs): layer.w13_weight_scale_from_checkpoint.data.fill_(2) layer.w2_weight_scale_from_checkpoint.data.fill_(2) + if replace_runtime_scales: + layer.w13_weight_scale = torch.nn.Parameter( + torch.zeros_like(layer.w13_weight_scale), requires_grad=False + ) + layer.w2_weight_scale = torch.nn.Parameter( + torch.zeros_like(layer.w2_weight_scale), requires_grad=False + ) fp8.process_weights_after_loading_mxfp8_moe(quant_method, layer) assert quant_method.moe_kernel is kernel @@ -870,8 +880,25 @@ def make_kernel(**kwargs): layer.w13_weight_scale, layer.w2_weight_scale, ) - assert tuple(id(parameter) for parameter in refit_parameters) == parameter_ids - assert tuple(parameter.data_ptr() for parameter in refit_parameters) == storage_ptrs + if replace_runtime_scales: + assert ( + tuple(id(parameter) for parameter in refit_parameters[:2]) + == parameter_ids[:2] + ) + assert ( + tuple(parameter.data_ptr() for parameter in refit_parameters[:2]) + == (storage_ptrs[:2]) + ) + assert quant_config.w1_scale is runtime_parameters[2] + assert quant_config.w2_scale is runtime_parameters[3] + assert torch.all(quant_config.w1_scale == 2) + assert torch.all(quant_config.w2_scale == 2) + else: + assert tuple(id(parameter) for parameter in refit_parameters) == parameter_ids + assert ( + tuple(parameter.data_ptr() for parameter in refit_parameters) + == storage_ptrs + ) assert all(torch.all(parameter == 2) for parameter in refit_parameters) assert kernel_calls[0] == { "moe_quant_config": quant_config, From b13119c1341c0934190eae85b5306a470c4a569c Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sun, 6 Sep 2026 22:56:33 -0700 Subject: [PATCH 74/76] test(refit): follow tuple weight iterator contract --- tests/unit/models/generation/test_vllm_fp8_quantization.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/unit/models/generation/test_vllm_fp8_quantization.py b/tests/unit/models/generation/test_vllm_fp8_quantization.py index 6766654ce61..70142f0a4af 100644 --- a/tests/unit/models/generation/test_vllm_fp8_quantization.py +++ b/tests/unit/models/generation/test_vllm_fp8_quantization.py @@ -1150,7 +1150,7 @@ def test_load_weights_accepts_prequantized_mxfp8_split_across_batches( fp8.load_weights([("model.weight_scale_from_checkpoint", scale)], model_runner) assert loaded == [ - ["model.weight", weight], + ("model.weight", weight), ("model.weight_scale_from_checkpoint", scale), ] From 0606ea87b011451f98c8f86a6c62e1c64706cdad Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sun, 6 Sep 2026 23:10:01 -0700 Subject: [PATCH 75/76] docs(refit): state prequantization lifetime contracts --- .../generation/vllm/quantization/fp8_train_utils.py | 9 ++++++++- nemo_rl/models/policy/workers/megatron_policy_worker.py | 1 + 2 files changed, 9 insertions(+), 1 deletion(-) diff --git a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py index d063026de5e..1dc3a38500a 100644 --- a/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py +++ b/nemo_rl/models/generation/vllm/quantization/fp8_train_utils.py @@ -125,7 +125,8 @@ def iter_mxfp8_prequantized_params( selected_names: Parameter names selected for MXFP8 prequantization. quantize_fn: MXFP8 quantization function. scratch_cache: Reusable stacking buffers keyed by device, dtype, and - CUDA stream. + CUDA stream. When omitted, reuse is limited to this export pass so + the stacking storage is released before training resumes. max_experts_per_batch: Maximum number of experts per quantization call. Yields: @@ -243,6 +244,8 @@ def flush_pending() -> Iterator[tuple[str, torch.Tensor]]: if source_stream is not None and source_stream != stack_stream: stack_stream.wait_stream(source_stream) tensor.record_stream(stack_stream) + # Stream objects must outlive this export pass; their raw CUDA + # handles form part of the scratch-storage identity. stream_id = ( int(stack_stream.cuda_stream) if stack_stream is not None else None ) @@ -265,6 +268,8 @@ def flush_pending() -> Iterator[tuple[str, torch.Tensor]]: scale_shape = (*first.shape[:-1], scale_columns) scale = scale.view(len(chunk), *scale_shape) for offset, index in enumerate(chunk_indices): + # These views avoid per-expert copies. Refit consumers must + # copy them before advancing past the current export batch. results[index] = value[offset], scale[offset], stack_stream for index, (_id, _proj, name, _tensor, _stream) in enumerate(pending): @@ -300,6 +305,8 @@ def flush_pending() -> Iterator[tuple[str, torch.Tensor]]: and len(pending_expert_ids) == max_experts_per_batch ): yield from flush_pending() + # Bridge export must hand off any private producer stream before yield; + # the ambient stream observed here defines tensor readiness. source_stream = ( torch.cuda.current_stream(tensor.device) if tensor.is_cuda else None ) diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 7c6c744ed69..65c43312590 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -2398,6 +2398,7 @@ def enable_refit_prequantize(self, param_names: list[str]) -> dict[str, Any]: def _maybe_prequantize_param( self, name: str, tensor: torch.Tensor ) -> Iterator[tuple[str, torch.Tensor]]: + """Single-tensor fallback; normal trainer export uses the batched iterator.""" if name not in self._refit_prequant_names: yield name, tensor return From b72f4d4e7730a62dca4e2e3c23332a2d9ca41fc7 Mon Sep 17 00:00:00 2001 From: seonjinn Date: Sun, 6 Sep 2026 23:42:46 -0700 Subject: [PATCH 76/76] test(refit): update offload call expectation Signed-off-by: seonjinn --- tests/unit/models/policy/test_megatron_worker.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index 10c4e552caa..acf6414304c 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -909,7 +909,7 @@ def test_offload_after_refit_routes_cleanup_by_mode( worker.offload_after_refit() worker.finalize_async_save.assert_called_once_with() - worker.move_model.assert_called_once_with(model, "cpu") + worker.move_model.assert_called_once_with(model, "cpu", move_params=True) model.eval.assert_called_once_with() if slim: worker._clear_fp8_caches.assert_called_once_with()