From cc0076a8f9b8d69873f17de5be975d520dd9d869 Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Tue, 30 Jun 2026 10:55:27 -0700 Subject: [PATCH 01/18] feat: add S3-based delta-compressed collective refit Signed-off-by: Hollow Man --- docs/design-docs/sparse-delta-refit.md | 35 + docs/index.md | 1 + examples/configs/grpo_math_1B.yaml | 5 + nemo_rl/algorithms/grpo.py | 64 +- nemo_rl/models/generation/__init__.py | 6 +- nemo_rl/models/generation/vllm/config.py | 18 + .../models/generation/vllm/vllm_backend.py | 713 +++++++++++++++++- .../models/generation/vllm/vllm_generation.py | 13 + nemo_rl/models/generation/vllm/vllm_worker.py | 292 ++++++- nemo_rl/models/policy/lm_policy.py | 41 + .../policy/workers/megatron_policy_worker.py | 58 ++ nemo_rl/utils/weight_transfer_s3.py | 122 +++ nemo_rl/utils/weight_transfer_s3_manifest.py | 430 +++++++++++ nemo_rl/utils/weight_transfer_sparse_codec.py | 245 ++++++ .../vllm_s3_sparse_weight_synchronizer.py | 122 +++ pyproject.toml | 2 + pyrefly.toml | 5 + .../models/generation/test_vllm_backend.py | 232 ++++++ .../models/generation/test_vllm_generation.py | 108 +++ .../unit/reference_configs/grpo_math_1B.yaml | 5 + uv.lock | 36 + 21 files changed, 2511 insertions(+), 42 deletions(-) create mode 100644 docs/design-docs/sparse-delta-refit.md create mode 100644 nemo_rl/utils/weight_transfer_s3.py create mode 100644 nemo_rl/utils/weight_transfer_s3_manifest.py create mode 100644 nemo_rl/utils/weight_transfer_sparse_codec.py create mode 100644 nemo_rl/weight_sync/vllm_s3_sparse_weight_synchronizer.py diff --git a/docs/design-docs/sparse-delta-refit.md b/docs/design-docs/sparse-delta-refit.md new file mode 100644 index 00000000000..0cfa9c2c687 --- /dev/null +++ b/docs/design-docs/sparse-delta-refit.md @@ -0,0 +1,35 @@ +# S3 Sparse-Delta vLLM Refit +For non-colocated Megatron policy workers and sync vLLM workers that share the +same checkpoint. Policy workers keep a CPU baseline, upload zstd-compressed +sparse deltas to S3, post receiver manifests, and commit the baseline only after +the global flush succeeds. + +## Config + +```yaml +backend: vllm +colocated: {enabled: false} +refit_transport: vllm_s3_sparse +delta_compression: + enabled: true + dtype: bf16 + sparse_bucket_size_bytes: 268435456 +vllm_cfg: + async_engine: false + expose_http_refit_server: true + http_refit_server_port: 8081 + http_refit_api_key_env_var: NRL_REFIT_API_KEY +``` + +S3 sparse refit requires `kv_cache_dtype: auto`; FP8 KV-cache scale sync is not +supported. Receiver tensors must have a direct QKV, MoE, Mamba, or generic TP +placement plan; transformed and FP8 weights fail before any delta is applied. +The receiver exposes `/nemo-rl/refit/s3-manifest`, with +`http_refit_api_key_env_var` auth when configured. Export chunks are capped by +`NRL_REFIT_S3_EXPORT_CHUNK_BYTES` and +`delta_compression.sparse_bucket_size_bytes`. Set `NRL_REFIT_S3_BUCKET` and, +when needed, `NRL_REFIT_S3_REGION` or `NRL_REFIT_S3_PREFIX`. AWS CRT performs +multipart transfer automatically; encode and end-to-end pipeline concurrency +default from available CPU cores and can be fixed with +`NRL_REFIT_S3_ENCODE_WORKERS` and `NRL_REFIT_S3_UPLOAD_WORKERS`. Track `REFIT_S3_TIMING`, +`REFIT_RECEIVER_TIMING`, and `REFIT_S3_GLOBAL_FLUSH` in cluster runs. diff --git a/docs/index.md b/docs/index.md index f57dc285104..4e86e4e47b9 100644 --- a/docs/index.md +++ b/docs/index.md @@ -322,6 +322,7 @@ design-docs/uv.md design-docs/dependency-management.md design-docs/chat-datasets.md design-docs/generation.md +design-docs/sparse-delta-refit.md design-docs/checkpointing.md design-docs/loss-functions.md design-docs/fsdp2-parallel-plan.md diff --git a/examples/configs/grpo_math_1B.yaml b/examples/configs/grpo_math_1B.yaml index e59d1675464..1a742b53cf3 100644 --- a/examples/configs/grpo_math_1B.yaml +++ b/examples/configs/grpo_math_1B.yaml @@ -339,6 +339,8 @@ policy: top_k: null stop_token_ids: null stop_strings: null + refit_transport: null # Set to "vllm_s3_sparse" to use S3 sparse-delta refit. + delta_compression: null # S3 sparse-delta refit config; null uses the existing refit path. mcore_generation_config: async_engine: false max_model_len: ${policy.max_total_sequence_length} # Engine-side max sequence length. @@ -376,6 +378,9 @@ policy: num_first_layers_in_bf16: 0 enable_vllm_metrics_logger: true # Set to true to enable vLLM internal metrics logger, turn off for better performance vllm_metrics_logger_interval: 0.5 # Interval in seconds to collect vLLM logger metrics + expose_http_refit_server: false # Start the internal sparse-delta refit endpoint on vLLM workers. + http_refit_api_key_env_var: null # Optional env var containing the internal refit API key. + http_refit_server_port: null # Optional fixed port for Kubernetes targetPorts. vllm_kwargs: {} colocated: # true: generation shares training GPUs diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index 6de9fd7a6e8..fe13717eb53 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -118,6 +118,9 @@ from nemo_rl.utils.nsys import maybe_gpu_profile_step from nemo_rl.utils.timer import TimeoutChecker, Timer from nemo_rl.utils.venvs import create_local_venv_on_each_node +from nemo_rl.weight_sync.vllm_s3_sparse_weight_synchronizer import ( + VllmS3SparseWeightSynchronizer, +) # =============================================================================== # Configuration @@ -907,6 +910,10 @@ def _spinup_nemo_gym(base_urls, model_name): # vllm model loading prefers clean environment, initialize policy_generation before policy in colocated mode backend = generation_config["backend"] generation_config["model_name"] = policy_config["model_name"] # Needed for vLLM + use_vllm_s3_sparse_refit = ( + backend == "vllm" + and generation_config.get("refit_transport") == "vllm_s3_sparse" + ) # Dictionary to store worker initialization timing stats for logging worker_init_timing_metrics = {} @@ -1095,6 +1102,30 @@ def initialize_generation_with_policy( elif backend == "vllm": # vLLM generation: setup config, then initialize with policy generation_config = cast(VllmConfig, generation_config) + vllm_cfg = generation_config["vllm_cfg"] + refit_transport = generation_config.get("refit_transport", None) + if refit_transport not in (None, "vllm_s3_sparse"): + raise ValueError(f"Unsupported vLLM refit transport {refit_transport!r}.") + if use_vllm_s3_sparse_refit: + delta_config = generation_config.get("delta_compression") + if ( + colocated_inference + or not policy_config["megatron_cfg"]["enabled"] + or vllm_cfg["async_engine"] + or vllm_cfg["precision"] == "fp8" + or vllm_cfg["kv_cache_dtype"].startswith("fp8") + or not delta_config + or not delta_config.get("enabled") + or generation_config.get("quant_cfg") + or generation_config.get("real_quant") + or not vllm_cfg.get("expose_http_refit_server") + ): + raise ValueError( + "vllm_s3_sparse requires a non-colocated Megatron policy, " + "synchronous BF16/FP16 vLLM, delta compression, an unquantized " + "rollout, and the refit HTTP server." + ) + if generation_config["vllm_cfg"]["precision"] == "fp8": assert loss_config.use_importance_sampling_correction, ( "Importance sampling must be enabled for vLLM FP8 generation for good convergence!" @@ -1232,7 +1263,7 @@ def init_vllm_then_policy(): policy.print_node_ip_and_gpu_id() # if it is not colocated inference, initialize collective communication for update weights - if not colocated_inference: + if not colocated_inference and not use_vllm_s3_sparse_refit: t0 = time.perf_counter() ip, port = train_cluster.get_master_address_and_port() print(f"Using ip: {ip}, port: {port} for collective communication", flush=True) @@ -1271,9 +1302,24 @@ def init_vllm_then_policy(): ray.get(futures_train + futures_inference) worker_init_timing_metrics["collective_init_time_s"] = time.perf_counter() - t0 - state_dict_info = policy.prepare_refit_info() - if policy_generation is not None: - policy_generation.prepare_refit_info(state_dict_info) + if use_vllm_s3_sparse_refit: + t0 = time.perf_counter() + assert isinstance(policy_generation, VllmGeneration) + policy_generation.weight_synchronizer = VllmS3SparseWeightSynchronizer( + policy, + policy_generation, + api_key_env_var=generation_config["vllm_cfg"].get( + "http_refit_api_key_env_var" + ), + ) + policy_generation.weight_synchronizer.init_communicator() + worker_init_timing_metrics["s3_sparse_refit_init_time_s"] = ( + time.perf_counter() - t0 + ) + else: + state_dict_info = policy.prepare_refit_info() + if policy_generation is not None: + policy_generation.prepare_refit_info(state_dict_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 @@ -2044,6 +2090,16 @@ def refit_policy_generation( timer: Optional Timer used to time the prepare/transfer/update phase kv_scales: Optional dictionary of KV cache scales for FP8 quantization. """ + if ( + isinstance(policy_generation, VllmGeneration) + and policy_generation.weight_synchronizer is not None + ): + policy_generation.weight_synchronizer.sync_weights( + timer=timer, + kv_scales=kv_scales, + ) + return + # Megatron generation backend needs explicit suspend/resume around refits. if isinstance(policy_generation, MegatronGeneration): policy_generation.suspend_for_refit() diff --git a/nemo_rl/models/generation/__init__.py b/nemo_rl/models/generation/__init__.py index 76756fc7176..26c82e50a2e 100644 --- a/nemo_rl/models/generation/__init__.py +++ b/nemo_rl/models/generation/__init__.py @@ -45,7 +45,11 @@ def configure_generation_config( if config["backend"] == "vllm": config = cast(VllmConfig, config) # set load_format - config["vllm_cfg"]["load_format"] = "auto" if is_eval else "dummy" + config["vllm_cfg"]["load_format"] = ( + "auto" + if is_eval or config.get("refit_transport") == "vllm_s3_sparse" + else "dummy" + ) speculative_config = config.get("vllm_kwargs", {}).get("speculative_config") if speculative_config and not is_eval and not has_refit_draft_weights: # Speculative decoding needs real draft weights at startup, since the diff --git a/nemo_rl/models/generation/vllm/config.py b/nemo_rl/models/generation/vllm/config.py index 44eb5d5d89c..1e14f572ae2 100644 --- a/nemo_rl/models/generation/vllm/config.py +++ b/nemo_rl/models/generation/vllm/config.py @@ -16,6 +16,8 @@ from nemo_rl.models.generation.interfaces import GenerationConfig +DeltaCompressionDType = Literal["fp16", "float16", "bf16", "bfloat16", "fp32", "float32"] # fmt: skip + class VllmSpecificArgs(TypedDict): tensor_parallel_size: int @@ -39,6 +41,12 @@ class VllmSpecificArgs(TypedDict): # Exposing vLLM as a server is useful in instances where the multi-turn rollout is performed with utilities outside of NeMo RL, but the user still wants to take advantage of the refit logic in NeMo RL that keeps the policy and generation up to date. # Currently it will expose the /tokenize and /v1/chat/completions endpoints. Later on we may expose /v1/completions or /v1/responses. expose_http_server: NotRequired[bool] + # Internal trusted endpoint for sparse delta refit payloads. + expose_http_refit_server: NotRequired[bool] + # Environment variable containing the internal refit API key. + http_refit_api_key_env_var: NotRequired[str | None] + # Fixed internal refit endpoint port for stable Kubernetes targetPorts. + http_refit_server_port: NotRequired[int | None] # These kwargs are passed to the vllm.LLM HTTP server Chat Completions endpoint config. Typically this will include things like tool parser, chat template, etc http_server_serving_chat_kwargs: NotRequired[dict[str, Any]] # Miscellaneous top level vLLM HTTP server arguments. @@ -52,9 +60,19 @@ class VllmSpecificArgs(TypedDict): reasoning_parser_plugin: NotRequired[str] +class VllmDeltaCompressionConfig(TypedDict): + enabled: bool + dtype: DeltaCompressionDType + sparse_bucket_size_bytes: int + baseline_mmap_dir: NotRequired[str | None] + + class VllmConfig(GenerationConfig): vllm_cfg: VllmSpecificArgs vllm_kwargs: NotRequired[dict[str, Any]] + # Null uses the existing NCCL refit; "vllm_s3_sparse" uses S3 sparse deltas. + refit_transport: NotRequired[Literal["vllm_s3_sparse"] | None] + delta_compression: NotRequired[VllmDeltaCompressionConfig | None] # quantization config quant_cfg: NotRequired[str | None] diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index f6b2e5adb94..f975c982d8b 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -12,9 +12,12 @@ # See the License for the specific language governing permissions and # limitations under the License. import gc +import io import re +import time import traceback -from typing import Any +from dataclasses import dataclass +from typing import Any, cast import torch import zmq @@ -24,9 +27,15 @@ calculate_aligned_size, rebuild_cuda_tensor_from_ipc, ) +from nemo_rl.utils import weight_transfer_sparse_codec as sparse_codec from nemo_rl.utils.nsys import wrap_with_nvtx_name from nemo_rl.utils.packed_tensor import packed_broadcast_consumer +_EXPERT_WEIGHT_RE = re.compile( + r"^(?P.*\.experts)\.(?P\d+)\." + r"(?Pgate_proj|up_proj|down_proj)\.weight$" +) + try: import vllm # noqa: F401 except ImportError: @@ -91,7 +100,29 @@ def _read_mtp_layer_weights_from_checkpoint( return weights +@dataclass(frozen=True) +class _SparseDeltaTargetPlan: + target: torch.Tensor | None + source_shape: tuple[int, ...] = () + source_strides: tuple[int, ...] = () + target_strides: tuple[int, ...] = () + target_offset: int = 0 + shard_dim: int | None = None + shard_start: int = 0 + shard_size: int = 0 + segment_shards: tuple[tuple[int, int, int], ...] = () + log_delta_transform: bool = False + identity: bool = False + + class VllmInternalWorkerExtension: + state_dict_info: dict[str, Any] | None = None + _direct_sparse_delta_targets: dict[str, torch.Tensor] | None = None + _direct_sparse_delta_modules: dict[str, torch.nn.Module] | None = None + _direct_sparse_delta_plan_cache: dict[str, _SparseDeltaTargetPlan | None] | None = ( + None + ) + def bind_numa(self) -> bool: """Pin this TP worker to its GPU's NUMA-local CPUs/memory. @@ -137,6 +168,12 @@ def report_device_id(self) -> str: return get_device_uuid(self.device.index) + def report_node_hostname(self) -> str: + """Return the host shared by worker processes on this node.""" + import socket + + return socket.gethostname() + def get_zmq_address(self): """Get the ZMQ address for the current device.""" return f"ipc:///tmp/{self.report_device_id()}.sock" @@ -165,9 +202,29 @@ def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: e.g. {tensor_name: (shape, dtype)} """ 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 + self._direct_sparse_delta_targets = None + self._direct_sparse_delta_modules = None + self._direct_sparse_delta_plan_cache = None + + def _process_weights_after_loading( + self, + model_config: Any, + target_device: torch.device, + ) -> None: + from vllm.config import set_current_vllm_config + from vllm.model_executor.model_loader.utils import ( + process_weights_after_loading, + ) + + with set_current_vllm_config(self.model_runner.vllm_config): + process_weights_after_loading( + self.model_runner.model, + model_config, + target_device, + ) def _maybe_process_fp8_kv_cache(self) -> None: - """Process weights after loading for FP8 KV cache (static scales).""" + """Process weights after loading for FP8 KV cache static scales.""" use_fp8_kv_cache = False if hasattr(self.model_runner.vllm_config, "cache_config"): kv_cache_dtype = getattr( @@ -176,27 +233,13 @@ def _maybe_process_fp8_kv_cache(self) -> None: use_fp8_kv_cache = ( kv_cache_dtype is not None and "fp8" in str(kv_cache_dtype).lower() ) - if not use_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, + self._process_weights_after_loading( + self.model_runner.model_config, + next(self.model_runner.model.parameters()).device, ) - # 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]], @@ -295,10 +338,12 @@ def load_mtp_weights_from_disk(self, model_path: str) -> bool: return False predictor = draft_model.model + mtp_start_layer_idx = cast(int, predictor.mtp_start_layer_idx) + num_mtp_layers = cast(int, predictor.num_mtp_layers) mtp_layer_indices = set( range( - predictor.mtp_start_layer_idx, - predictor.mtp_start_layer_idx + predictor.num_mtp_layers, + mtp_start_layer_idx, + mtp_start_layer_idx + num_mtp_layers, ) ) weights = _read_mtp_layer_weights_from_checkpoint(model_path, mtp_layer_indices) @@ -353,6 +398,529 @@ def _load_weights(self, weights): self._load_draft_weights(draft_weights) + def _apply_sparse_weight_deltas( + self, + payload_tensors: tuple[torch.Tensor, torch.Tensor], + metadata: list[dict[str, Any]], + ) -> None: + """Apply sparse deltas directly after validating every target plan.""" + if self._direct_sparse_delta_uses_loader_transform(): + raise RuntimeError( + "Direct sparse delta refit does not support transformed or FP8 weights." + ) + + targets = self._direct_sparse_delta_target_map() + raw_locations, raw_values = payload_tensors + plans = [] + for item in metadata: + plan = self._cached_direct_sparse_delta_target_plan(item, targets) + if plan is None: + raise RuntimeError( + f"No direct sparse delta target plan for {item['name']!r}." + ) + plans.append((item, plan)) + + with torch.no_grad(): + for item, plan in plans: + target = plan.target + if target is None: + continue + + value_start = int(item["value_start"]) + value_end = int(item["value_end"]) + values = raw_values[value_start:value_end].to( + device=target.device, + dtype=target.dtype, + non_blocking=True, + ) + if plan.identity and item["index_encoding"] == "range": + range_start = int(item["range_start"]) + range_count = value_end - value_start + target.data.view(-1).narrow(0, range_start, range_count).add_( + values + ) + continue + + locations = sparse_codec.sparse_locations_for_item( + item, + raw_locations, + device=target.device, + ) + locations, values = self._local_sparse_delta_update_inputs( + locations, + values, + plan, + ) + if locations.numel(): + target_flat = target.data.view(-1) + if plan.log_delta_transform: + current = target_flat.index_select(0, locations) + updated = current * values.float().exp().to(dtype=current.dtype) + target_flat.index_copy_(0, locations, updated) + else: + target_flat.index_add_(0, locations, values) + + def _direct_sparse_delta_target_map(self) -> dict[str, torch.Tensor]: + if self._direct_sparse_delta_targets is None: + self._direct_sparse_delta_targets = dict( + self.model_runner.model.named_parameters() + ) + self._direct_sparse_delta_targets.update( + self.model_runner.model.named_buffers() + ) + return self._direct_sparse_delta_targets + + def _direct_sparse_delta_named_module( + self, + module_name: str, + ) -> torch.nn.Module | None: + if self._direct_sparse_delta_modules is None: + self._direct_sparse_delta_modules = dict( + self.model_runner.model.named_modules() + ) + return self._direct_sparse_delta_modules.get(module_name) + + def _direct_sparse_delta_uses_loader_transform(self) -> bool: + architectures = self.model_runner.vllm_config.model_config.architectures + if {"GptOssForCausalLM", "Gemma3ForConditionalGeneration"} & set(architectures): + return True + + from nemo_rl.models.generation.vllm.quantization import fp8 + + return fp8.is_fp8_model(self.model_runner.vllm_config) + + def _cached_direct_sparse_delta_target_plan( + self, + item: dict[str, Any], + targets: dict[str, torch.Tensor], + ) -> _SparseDeltaTargetPlan | None: + name = str(item["name"]) + if self._direct_sparse_delta_plan_cache is None: + self._direct_sparse_delta_plan_cache = {} + if name not in self._direct_sparse_delta_plan_cache: + self._direct_sparse_delta_plan_cache[name] = ( + self._direct_sparse_delta_target_plan(item, targets) + ) + return self._direct_sparse_delta_plan_cache[name] + + def _direct_sparse_delta_target_plan( + self, + item: dict[str, Any], + targets: dict[str, torch.Tensor], + ) -> _SparseDeltaTargetPlan | None: + name = str(item["name"]) + if name.startswith("mtp."): + return _SparseDeltaTargetPlan(target=None) + target_name = self._map_direct_sparse_delta_name(name) + if target_name is None: + return None + if ".mixer." in target_name: + mamba_plan = self._direct_sparse_delta_mamba2_plan( + item, target_name, targets + ) + if mamba_plan is not None: + return mamba_plan + if ".mixer.conv1d." in target_name or ".mixer.in_proj." in target_name: + return None + if any(f".{candidate}_proj." in target_name for candidate in ("q", "k", "v")): + return self._direct_sparse_delta_qkv_plan(item, target_name, targets) + if _EXPERT_WEIGHT_RE.match(target_name): + return self._direct_sparse_delta_expert_plan(item, target_name, targets) + + target = targets.get(target_name) + if target is None: + return None + source_shape = tuple(item["shape"]) + target_shape = tuple(target.shape) + if target_shape == source_shape: + return self._make_sparse_delta_target_plan(target, source_shape) + return self._direct_sparse_delta_shard_plan(item, target) + + def _direct_sparse_delta_qkv_plan( + self, + item: dict[str, Any], + target_name: str, + targets: dict[str, torch.Tensor], + ) -> _SparseDeltaTargetPlan | None: + shard_id = next(x for x in "qkv" if f".{x}_proj." in target_name) + packed_name = target_name.replace(f".{shard_id}_proj.", ".qkv_proj.", 1) + target = targets.get(packed_name) + if target is None: + return None + output_dim = int(cast(Any, target).output_dim) % target.ndim + module = cast( + Any, + getattr(getattr(target, "weight_loader", None), "__self__", None) + or self._direct_sparse_delta_named_module(packed_name.rsplit(".", 1)[0]), + ) + shard_offset = int(module._get_shard_offset_mapping(shard_id)) + shard_size = int(module._get_shard_size_mapping(shard_id)) + shard_rank = int(module.tp_rank) + if shard_id != "q": + shard_rank //= int(module.num_kv_head_replicas) + + source_shape = tuple(item["shape"]) + shard_start = shard_rank * shard_size + if source_shape[output_dim] < shard_start: + return _SparseDeltaTargetPlan(target=None) + return self._make_sparse_delta_target_plan( + target, + source_shape=source_shape, + shard_dim=output_dim, + shard_start=shard_start, + shard_size=min(shard_size, source_shape[output_dim] - shard_start), + target_offset=shard_offset * target.stride(output_dim), + ) + + def _direct_sparse_delta_mamba2_plan( + self, + item: dict[str, Any], + target_name: str, + targets: dict[str, torch.Tensor], + ) -> _SparseDeltaTargetPlan | None: + target = targets.get(target_name) + if target is None: + return None + + if target_name.endswith(".A"): + source_shape = tuple(item["shape"]) + if tuple(target.shape) == source_shape: + return self._make_sparse_delta_target_plan( + target, + source_shape=source_shape, + log_delta_transform=True, + ) + return self._direct_sparse_delta_shard_plan( + item, + target, + log_delta_transform=True, + ) + + if not (".mixer.conv1d." in target_name or ".mixer.in_proj." in target_name): + return None + + mixer_name = target_name.split(".mixer.", 1)[0] + ".mixer" + attrs = cast(Any, self._direct_sparse_delta_named_module(mixer_name)) + tp_size = int(attrs.tp_size) + if tp_size <= 1: + return None + intermediate_size = int(attrs.intermediate_size) + groups_ssm_state_size = int(attrs.groups_ssm_state_size) + num_heads = int(attrs.num_heads) + source_shape = tuple(item["shape"]) + fixed_size = intermediate_size + if ".mixer.in_proj." in target_name: + fixed_size += intermediate_size + num_heads + group_size, remainder = divmod(source_shape[0] - fixed_size, 2) + extra_group_size = groups_ssm_state_size - group_size + if remainder or group_size <= 0 or extra_group_size < 0: + return None + tp_rank = int( + getattr(target, "tp_rank", self._direct_sparse_delta_tp_rank(tp_size)) + ) + intermediate = (intermediate_size, 0, False) + group = (groups_ssm_state_size, extra_group_size, extra_group_size > 0) + segment_specs = ( + (intermediate, group, group) + if ".mixer.conv1d." in target_name + else (intermediate, intermediate, group, group, (num_heads, 0, False)) + ) + + target_shape = tuple(target.shape) + source_to_target_dims = tuple(range(len(source_shape))) + if len(target_shape) == len(source_shape) + 1 and target_shape[1] == 1: + source_to_target_dims = (0, *range(2, len(target_shape))) + elif len(target_shape) != len(source_shape): + return None + segment_shards: list[tuple[int, int, int]] = [] + target_start = 0 + source_start = 0 + for full_dim, extra, duplicate_groups in segment_specs: + shard_size = full_dim // tp_size + rank = 0 if duplicate_groups else tp_rank + source_dim = full_dim - extra + source_local_start = source_start + rank * shard_size + take = min(shard_size, source_dim - rank * shard_size) + if take > 0: + segment_shards.append((source_local_start, target_start, take)) + target_start += shard_size + source_start += source_dim + if source_shape[0] != source_start or target_shape[0] != target_start: + return None + + return self._make_sparse_delta_target_plan( + target, + source_shape=source_shape, + source_to_target_dims=source_to_target_dims, + shard_dim=0, + segment_shards=tuple(segment_shards), + ) + + def _direct_sparse_delta_expert_plan( + self, + item: dict[str, Any], + target_name: str, + targets: dict[str, torch.Tensor], + ) -> _SparseDeltaTargetPlan | None: + match = cast(re.Match[str], _EXPERT_WEIGHT_RE.match(target_name)) + + prefix = match.group("prefix") + global_expert_id = int(match.group("expert")) + proj = match.group("proj") + packed_weight, shard_id = { + "gate_proj": ("w13_weight", "w1"), + "up_proj": ("w13_weight", "w3"), + "down_proj": ("w2_weight", "w2"), + }[proj] + packed_name = f"{prefix}.{packed_weight}" + + target = targets.get(packed_name) + if target is None: + return None + module_attrs = cast( + Any, + getattr(getattr(target, "weight_loader", None), "__self__", None) + or self._direct_sparse_delta_named_module(packed_name.rsplit(".", 1)[0]), + ) + if shard_id == "w3" and not module_attrs.moe_config.is_act_and_mul: + shard_id = "w1" + local_expert_id = int( + module_attrs._map_global_expert_id_to_local_expert_id(global_expert_id) + ) + if local_expert_id < 0: + return _SparseDeltaTargetPlan(target=None) + + source_shape = tuple(item["shape"]) + target_shape = tuple(target.shape) + shard_dim = 1 if shard_id == "w2" else 0 + if local_expert_id >= target_shape[0]: + return None + target_shard_dim = shard_dim + 1 + + shard_size = target_shape[target_shard_dim] + if shard_id in ("w1", "w3") and module_attrs.moe_config.is_act_and_mul: + shard_size //= 2 + target_shard_offset = shard_size if shard_id == "w3" else 0 + if target_shape[target_shard_dim] < target_shard_offset + shard_size: + return None + tp_rank = int(module_attrs.tp_rank) + shard_start = tp_rank * shard_size + if source_shape[shard_dim] < shard_start: + return _SparseDeltaTargetPlan(target=None) + + return self._make_sparse_delta_target_plan( + target, + source_shape=source_shape, + source_to_target_dims=tuple(dim + 1 for dim in range(len(source_shape))), + target_offset=( + local_expert_id * target.stride(0) + + target_shard_offset * target.stride(target_shard_dim) + ), + shard_dim=shard_dim, + shard_start=shard_start, + shard_size=min(shard_size, source_shape[shard_dim] - shard_start), + ) + + def _make_sparse_delta_target_plan( + self, + target: torch.Tensor, + source_shape: tuple[int, ...], + *, + source_to_target_dims: tuple[int, ...] | None = None, + target_offset: int = 0, + shard_dim: int | None = None, + shard_start: int = 0, + shard_size: int = 0, + segment_shards: tuple[tuple[int, int, int], ...] = (), + log_delta_transform: bool = False, + ) -> _SparseDeltaTargetPlan | None: + if source_to_target_dims is None: + source_to_target_dims = tuple(range(len(source_shape))) + target_shape = tuple(target.shape) + ignored_dim = ( + shard_dim if shard_dim is not None else 0 if segment_shards else -1 + ) + if len(source_to_target_dims) != len(source_shape) or any( + target_dim >= len(target_shape) + or ( + source_dim != ignored_dim + and source_shape[source_dim] != target_shape[target_dim] + ) + for source_dim, target_dim in enumerate(source_to_target_dims) + ): + return None + identity = ( + shard_dim is None + and target_offset == 0 + and not segment_shards + and not log_delta_transform + and source_to_target_dims == tuple(range(len(source_shape))) + and source_shape == target_shape + ) + return _SparseDeltaTargetPlan( + target=target, + source_shape=source_shape, + source_strides=torch.empty(source_shape, device="meta").stride(), + target_strides=tuple( + target.stride(target_dim) for target_dim in source_to_target_dims + ), + target_offset=target_offset, + shard_dim=shard_dim, + shard_start=shard_start, + shard_size=shard_size, + segment_shards=segment_shards, + log_delta_transform=log_delta_transform, + identity=identity, + ) + + def _direct_sparse_delta_shard_plan( + self, + item: dict[str, Any], + target: torch.Tensor, + *, + log_delta_transform: bool = False, + ) -> _SparseDeltaTargetPlan | None: + source_shape = tuple(item["shape"]) + target_shape = tuple(target.shape) + if len(source_shape) != len(target_shape): + return None + + candidate_dims = list( + dict.fromkeys( + dim % len(source_shape) + for attr in ("output_dim", "input_dim") + if isinstance(dim := getattr(target, attr, None), int) + ) + ) + if not candidate_dims: + candidate_dims = [ + dim + for dim, (source_dim, target_dim) in enumerate( + zip(source_shape, target_shape, strict=True) + ) + if source_dim != target_dim + ] + if len(candidate_dims) != 1: + return None + + for shard_dim in candidate_dims: + shard_size = target_shape[shard_dim] + tp_size = int(getattr(target, "tp_size", 1)) + if tp_size <= 1: + if shard_size <= 0 or source_shape[shard_dim] % shard_size: + continue + tp_size = source_shape[shard_dim] // shard_size + if tp_size <= 1: + continue + if source_shape[shard_dim] > shard_size * tp_size: + continue + tp_rank = int( + getattr(target, "tp_rank", self._direct_sparse_delta_tp_rank(tp_size)) + ) + plan = self._make_sparse_delta_target_plan( + target=target, + source_shape=source_shape, + shard_dim=shard_dim, + shard_start=tp_rank * shard_size, + shard_size=shard_size, + log_delta_transform=log_delta_transform, + ) + if plan is not None: + return plan + return None + + def _local_sparse_delta_update_inputs( + self, + locations: torch.Tensor, + values: torch.Tensor, + plan: _SparseDeltaTargetPlan, + ) -> tuple[torch.Tensor, torch.Tensor]: + if plan.identity: + return locations, values + + source_shape = plan.source_shape + source_strides = plan.source_strides + target_strides = plan.target_strides + shard_dim = plan.shard_dim + + if source_strides == target_strides: + if shard_dim is None: + return locations + plan.target_offset, values + if shard_dim == 0: + shard_stride = source_strides[0] + shard_coords = torch.div(locations, shard_stride, rounding_mode="floor") + if plan.segment_shards: + mapped_locations = locations + plan.target_offset + keep = torch.zeros_like(locations, dtype=torch.bool) + for source_start, target_start, take in plan.segment_shards: + segment = (shard_coords >= source_start) & ( + shard_coords < source_start + take + ) + mapped_locations[segment] += ( + target_start - source_start + ) * shard_stride + keep |= segment + return mapped_locations[keep], values[keep] + shard_end = min( + plan.shard_start + plan.shard_size, source_shape[shard_dim] + ) + keep = (shard_coords >= plan.shard_start) & (shard_coords < shard_end) + return ( + locations[keep] + + plan.target_offset + - plan.shard_start * shard_stride, + values[keep], + ) + + selected_locations = locations + selected_values = values + + if shard_dim is not None: + shard_coords = torch.div( + locations, + source_strides[shard_dim], + rounding_mode="floor", + ).remainder(source_shape[shard_dim]) + shard_end = min( + plan.shard_start + plan.shard_size, + source_shape[shard_dim], + ) + keep = (shard_coords >= plan.shard_start) & (shard_coords < shard_end) + selected_locations = locations[keep] + selected_values = values[keep] + if selected_locations.numel() == 0: + return selected_locations, selected_values + + local_locations = torch.full_like(selected_locations, plan.target_offset) + for dim, (source_stride, target_stride) in enumerate( + zip(source_strides, target_strides, strict=True) + ): + coord = torch.div( + selected_locations, + source_stride, + rounding_mode="floor", + ).remainder(source_shape[dim]) + if dim == plan.shard_dim: + coord = coord - plan.shard_start + local_locations.add_(coord * target_stride) + return local_locations, selected_values + + def _direct_sparse_delta_tp_rank(self, tp_size: int) -> int: + if tp_size <= 1: + return 0 + rank = int(getattr(self, "rank", 0)) + return rank % tp_size + + def _map_direct_sparse_delta_name(self, name: str) -> str | None: + mapper = getattr(self.model_runner.model, "hf_to_vllm_mapper", None) + if mapper is not None: + name = cast(Any, mapper)._map_name(name) + if name is None: + return None + if name.startswith("draft."): + return None + return name + @wrap_with_nvtx_name("vllm_internal_worker_extension/update_weights_via_ipc_zmq") def update_weights_via_ipc_zmq(self) -> bool: """Receive and update model weights via ZMQ IPC socket. @@ -371,26 +939,24 @@ def update_weights_via_ipc_zmq(self) -> bool: if payload == IPCProtocol.COMPLETE: # means the update is done - from vllm.config import set_current_vllm_config - from vllm.model_executor.model_loader.utils import ( - process_weights_after_loading, - ) - - with set_current_vllm_config(self.model_runner.vllm_config): - process_weights_after_loading( - self.model_runner.model, self.model_config, self.device - ) + self._process_weights_after_loading(self.model_config, self.device) self.zmq_socket.send(IPCProtocol.ACK.value.encode()) break ipc_handle, list_keys, used_bytes = payload buffer = rebuild_cuda_tensor_from_ipc(ipc_handle, self.device.index) + state_dict_info = self.state_dict_info + if state_dict_info is None: + raise RuntimeError( + "state_dict_info is not prepared. " + "Call prepare_refit_info before loading weights." + ) weight = None weights = [] offset = 0 for key in list_keys: - shape, dtype = self.state_dict_info[key] # pyrefly + shape, dtype = state_dict_info[key] if isinstance(shape, list): shape = torch.Size(shape) @@ -482,6 +1048,91 @@ def update_weights_from_collective(self) -> bool: torch.cuda.empty_cache() return True + @wrap_with_nvtx_name( + "vllm_internal_worker_extension/update_weights_from_serialized_sparse_payload" + ) + def update_weights_from_serialized_sparse_payload( + self, + serialized_payload: bytes, + synchronize: bool = True, + ) -> dict[str, Any]: + """Apply one serialized sparse-delta payload received from S3.""" + return self._load_and_apply_sparse_payload( + io.BytesIO(serialized_payload), synchronize + ) + + def _load_and_apply_sparse_payload( + self, + source: str | io.BytesIO, + synchronize: bool, + ) -> dict[str, Any]: + started = time.perf_counter() + payload = cast( + sparse_codec.TensorPayload, + torch.load( + source, + map_location="cpu", + weights_only=True, + ), + ) + deserialize_s = time.perf_counter() - started + result = self._apply_sparse_request(payload, synchronize=synchronize) + result["receiver_deserialize_s"] = deserialize_s + result["receiver_total_s"] = time.perf_counter() - started + return result + + @wrap_with_nvtx_name( + "vllm_internal_worker_extension/update_weights_from_sparse_payload_files" + ) + def update_weights_from_sparse_payload_files( + self, + *payload_paths: str, + synchronize: bool = True, + ) -> dict[str, Any]: + """Apply sparse payloads in FIFO order with one final sync.""" + if not payload_paths: + raise ValueError("A sparse refit batch must contain at least one payload.") + started = time.perf_counter() + deserialize_s = 0.0 + sparse_apply_s = 0.0 + for path in payload_paths: + result = self._load_and_apply_sparse_payload(path, synchronize=False) + deserialize_s += float(result["receiver_deserialize_s"]) + sparse_apply_s += float(result["receiver_sparse_apply_s"]) + if synchronize and torch.cuda.is_available(): + torch.cuda.synchronize(self.device) + return { + "ok": True, + "receiver_deserialize_s": deserialize_s, + "receiver_sparse_apply_s": sparse_apply_s, + "receiver_total_s": time.perf_counter() - started, + } + + def _apply_sparse_request( + self, + payload: sparse_codec.TensorPayload, + *, + synchronize: bool, + ) -> dict[str, Any]: + locations, values, metadata = payload + + sparse_started = time.perf_counter() + self._apply_sparse_weight_deltas((locations, values), metadata) + sparse_apply_s = time.perf_counter() - sparse_started + + if synchronize and torch.cuda.is_available(): + torch.cuda.synchronize(self.device) + return { + "ok": True, + "receiver_sparse_apply_s": sparse_apply_s, + } + + def synchronize_device(self) -> dict[str, Any]: + """Synchronize this vLLM worker's CUDA device after deferred refit applies.""" + if torch.cuda.is_available(): + torch.cuda.synchronize(self.device) + return {"ok": True} + def cleanup(self) -> None: """Shutdown and cleanup resources.""" # Close ZMQ socket and context if they exist diff --git a/nemo_rl/models/generation/vllm/vllm_generation.py b/nemo_rl/models/generation/vllm/vllm_generation.py index d68bf512bcc..058347c9f66 100644 --- a/nemo_rl/models/generation/vllm/vllm_generation.py +++ b/nemo_rl/models/generation/vllm/vllm_generation.py @@ -43,6 +43,7 @@ compute_spec_decode_metrics, resolve_generation_worker_cls, ) +from nemo_rl.weight_sync.interfaces import WeightSynchronizer logger = logging.getLogger(__name__) @@ -102,6 +103,7 @@ def __init__( # Store config self.cfg = config self._defer_model_load = defer_model_load + self.weight_synchronizer: WeightSynchronizer | None = None self.tp_size = self.cfg["vllm_cfg"]["tensor_parallel_size"] self.pp_size = self.cfg["vllm_cfg"]["pipeline_parallel_size"] self.ep_size = self.cfg["vllm_cfg"]["expert_parallel_size"] @@ -928,6 +930,17 @@ def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: # Wait for all futures to complete ray.get(futures) + def report_refit_server_base_urls(self) -> list[str]: + """Return base URLs for vLLM workers exposing sparse refit endpoints.""" + if not self.worker_group or not self.worker_group.workers: + raise RuntimeError("Worker group is not initialized") + + futures = self.worker_group.run_all_workers_single_data( + "report_refit_server_base_url", + run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], + ) + return [url for url in ray.get(futures) if url] + def update_weights_via_ipc_zmq(self) -> list[ray.ObjectRef]: """Update weights of the policy using IPC handles via ZMQ socket.""" if not self.worker_group or not self.worker_group.workers: diff --git a/nemo_rl/models/generation/vllm/vllm_worker.py b/nemo_rl/models/generation/vllm/vllm_worker.py index 8329da0b93c..3518338b671 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker.py +++ b/nemo_rl/models/generation/vllm/vllm_worker.py @@ -12,11 +12,18 @@ # See the License for the specific language governing permissions and # limitations under the License. +import asyncio import copy import gc import logging import os import sys +import tempfile +import threading +import time +import traceback +from collections import deque +from concurrent.futures import Future, ThreadPoolExecutor from typing import Any, Optional, cast import ray @@ -25,8 +32,12 @@ from nemo_rl.distributed.batched_data_dict import BatchedDataDict from nemo_rl.distributed.virtual_cluster import ( + DEFAULT_GENERATION_PORT_RANGE_HIGH, + DEFAULT_GENERATION_PORT_RANGE_LOW, DEFAULT_VLLM_PORT_RANGE_LOW, DEFAULT_VLLM_PORTS_PER_ENGINE, + _get_free_port_local, + _get_node_ip_local, ) from nemo_rl.distributed.worker_group_utils import get_nsight_config_if_pattern_matches from nemo_rl.models.generation.interfaces import ( @@ -47,6 +58,18 @@ from nemo_rl.models.policy.utils import is_vllm_v1_engine_enabled from nemo_rl.utils.nsys import wrap_with_nvtx_name from nemo_rl.utils.nvml import log_gpu_memory_diagnostics +from nemo_rl.utils.weight_transfer_s3_manifest import ( + G_VLLM_REFIT_API_KEY_HEADER, + G_VLLM_REFIT_FLUSH_PATH, + G_VLLM_REFIT_S3_MANIFEST_PATH, + download_s3_refit_payload, + merge_vllm_refit_receiver_timing, + vllm_refit_api_key, +) + +G_REFIT_APPLY_QUEUE_DEPTH_ENV = "NRL_REFIT_APPLY_QUEUE_DEPTH" +G_REFIT_APPLY_BATCH_SIZE_ENV = "NRL_REFIT_APPLY_BATCH_SIZE" +G_REFIT_BATCH_STAGING_DIR_ENV = "NRL_REFIT_BATCH_STAGING_DIR" logger = logging.getLogger(__name__) @@ -244,6 +267,28 @@ def __init__( if bundle_indices is not None and len(bundle_indices) == 1: bind_to_gpu_numa(int(ray.get_gpu_ids()[0])) + self._refit_apply_queue_lock = threading.Lock() + self._refit_apply_executor = ThreadPoolExecutor(max_workers=1) + self._refit_apply_futures: deque[Future[dict[str, Any]]] = deque() + self._refit_apply_pending_payloads: list[bytes] = [] + self._refit_apply_payload_count = 0 + self._refit_apply_batch_count = 0 + self._refit_workers_share_node = False + self._refit_apply_queue_depth = int( + os.getenv(G_REFIT_APPLY_QUEUE_DEPTH_ENV) or 2 + ) + self._refit_apply_batch_size = int(os.getenv(G_REFIT_APPLY_BATCH_SIZE_ENV) or 8) + self._refit_batch_staging_dir = ( + os.getenv(G_REFIT_BATCH_STAGING_DIR_ENV) or "/dev/shm" + ) + if self._refit_apply_queue_depth < 1: + raise ValueError(f"{G_REFIT_APPLY_QUEUE_DEPTH_ENV} must be >= 1.") + if self._refit_apply_batch_size < 1: + raise ValueError(f"{G_REFIT_APPLY_BATCH_SIZE_ENV} must be >= 1.") + self.refit_server_base_url: str | None = None + self.refit_server: Any | None = None + self.refit_server_thread: threading.Thread | None = None + self._init_config( config, bundle_indices, fraction_of_gpus, seed, extra_env_vars ) @@ -650,6 +695,155 @@ def _get_raw_spec_counters(self) -> dict[str, float | list[float]]: metrics[metric.name] = metric.value return metrics + def _enqueue_sparse_payload_apply( + self, + payload: bytes, + ) -> dict[str, Any]: + completed: list[Future[dict[str, Any]]] = [] + with self._refit_apply_queue_lock: + while self._refit_apply_futures and ( + self._refit_apply_futures[0].done() + or len(self._refit_apply_futures) >= self._refit_apply_queue_depth + ): + completed.append(self._refit_apply_futures.popleft()) + response = self._collect_refit_apply_results(completed) + self._refit_apply_pending_payloads.append(payload) + self._refit_apply_payload_count += 1 + if len(self._refit_apply_pending_payloads) == self._refit_apply_batch_size: + self._submit_pending_sparse_payloads() + return response + + def _submit_pending_sparse_payloads(self) -> None: + payloads = tuple(self._refit_apply_pending_payloads) + self._refit_apply_pending_payloads.clear() + self._refit_apply_futures.append( + self._refit_apply_executor.submit( + self.update_weights_from_serialized_sparse_payloads, + payloads, + False, + ) + ) + self._refit_apply_batch_count += 1 + + def _collect_refit_apply_results( + self, + futures: list[Future[dict[str, Any]]], + *, + synchronize: bool = False, + ) -> dict[str, Any]: + results = [future.result() for future in futures] + if synchronize and results: + assert self.llm is not None + self.llm.collective_rpc("synchronize_device", args=()) + timing: dict[str, float] = {} + merge_vllm_refit_receiver_timing(timing, results, maximum=False) + return { + "ok": True, + "payloads": sum(int(result.get("payloads", 0)) for result in results), + **timing, + } + + @staticmethod + def _refit_collective_response(worker_results: Any) -> dict[str, Any]: + return { + "ok": True, + **merge_vllm_refit_receiver_timing( + {}, cast(list[dict[str, Any]], worker_results), maximum=True + ), + } + + def _flush_queued_sparse_payloads(self) -> dict[str, Any]: + started = time.perf_counter() + with self._refit_apply_queue_lock: + if self._refit_apply_pending_payloads: + self._submit_pending_sparse_payloads() + futures = list(self._refit_apply_futures) + self._refit_apply_futures.clear() + payload_count = self._refit_apply_payload_count + batch_count = self._refit_apply_batch_count + self._refit_apply_payload_count = 0 + self._refit_apply_batch_count = 0 + response = self._collect_refit_apply_results(futures, synchronize=True) + response.update( + payloads=payload_count, + batches=batch_count, + seconds=time.perf_counter() - started, + ) + if futures: + print( + "REFIT_RECEIVER_TIMING " + f"payloads={payload_count} batches={batch_count} " + f"total_s={response['seconds']:.3f} " + f"payload_total_s={response.get('receiver_total_s', 0.0):.3f}", + flush=True, + ) + return response + + def update_weights_from_serialized_sparse_payload( + self, serialized_payload: bytes, synchronize: bool = True + ) -> dict[str, Any]: + return self.update_weights_from_serialized_sparse_payloads( + (serialized_payload,), synchronize + ) + + def update_weights_from_serialized_sparse_payloads( + self, + serialized_payloads: tuple[bytes, ...], + synchronize: bool = True, + ) -> dict[str, Any]: + raise NotImplementedError + + async def _apply_s3_manifest_payload( + self, + manifest: dict[str, Any], + ) -> dict[str, Any]: + started = time.perf_counter() + body = await asyncio.to_thread(download_s3_refit_payload, manifest) + download_s = time.perf_counter() - started + result = await asyncio.to_thread(self._enqueue_sparse_payload_apply, body) + result.update(payloads=1, receiver_s3_download_s=download_s) + return result + + def _setup_vllm_refit_api_server(self, app: Any) -> None: + from fastapi import Request + from fastapi.responses import JSONResponse + + token = vllm_refit_api_key( + self.cfg["vllm_cfg"].get("http_refit_api_key_env_var") + ) + + async def respond(raw_request: Request, *, flush: bool = False) -> JSONResponse: + if ( + token is not None + and raw_request.headers.get(G_VLLM_REFIT_API_KEY_HEADER) != token + ): + return JSONResponse( + content={"ok": False, "error": "unauthorized"}, status_code=403 + ) + try: + result = ( + await asyncio.to_thread(self._flush_queued_sparse_payloads) + if flush + else await self._apply_s3_manifest_payload(await raw_request.json()) + ) + except Exception as exc: + result = {"ok": False, "error": str(exc)} + return JSONResponse( + content=result, + status_code=200 if result.get("ok") is True else 500, + ) + + @app.post(G_VLLM_REFIT_S3_MANIFEST_PATH) + async def apply_s3_manifest_refit(raw_request: Request) -> JSONResponse: + return await respond(raw_request) + + @app.post(G_VLLM_REFIT_FLUSH_PATH) + async def flush_sparse_delta_refit(raw_request: Request) -> JSONResponse: + return await respond(raw_request, flush=True) + + def report_refit_server_base_url(self) -> str | None: + return self.refit_server_base_url + class VllmGenerationWorkerImpl(BaseVllmGenerationWorker): def _create_engine(self, llm_kwargs: dict[str, Any]) -> None: @@ -665,6 +859,38 @@ def post_init(self): self.llm.collective_rpc( "load_mtp_weights_from_disk", args=(self.model_name,) ) + if self.cfg["vllm_cfg"].get("expose_http_refit_server"): + self._refit_workers_share_node = ( + len(set(self.llm.collective_rpc("report_node_hostname", args=()))) == 1 + ) + self._setup_vllm_refit_server() + + def _setup_vllm_refit_server(self) -> None: + import uvicorn + from fastapi import FastAPI + + app = FastAPI() + self._setup_vllm_refit_api_server(app) + port = self.cfg["vllm_cfg"].get( + "http_refit_server_port" + ) or _get_free_port_local( + self.cfg.get("port_range_low", DEFAULT_GENERATION_PORT_RANGE_LOW), + self.cfg.get("port_range_high", DEFAULT_GENERATION_PORT_RANGE_HIGH), + ) + server = uvicorn.Server( + uvicorn.Config( + app, + host="0.0.0.0", + port=port, + timeout_keep_alive=120, + ) + ) + thread = threading.Thread(target=server.run, daemon=True) + thread.start() + self.refit_server_base_url = f"http://{_get_node_ip_local()}:{port}" + self.refit_server = server + self.refit_server_thread = thread + print(f"Starting vLLM refit server on {self.refit_server_base_url}", flush=True) def init_collective( self, @@ -790,8 +1016,6 @@ def generate( 1 ].logprob except Exception: - import traceback - traceback.print_exc() logprobs_list.append(full_logprobs) @@ -989,8 +1213,6 @@ def update_weights_via_ipc_zmq(self) -> bool: return True except Exception as e: print(f"Exception during collective_rpc for weight update: {e}") - import traceback - traceback.print_exc() return False @@ -1020,11 +1242,60 @@ def update_weights_from_collective(self) -> bool: return True except Exception as e: print(f"Exception during collective_rpc for weight update: {e}") - import traceback - traceback.print_exc() return False + def update_weights_from_serialized_sparse_payloads( + self, + serialized_payloads: tuple[bytes, ...], + synchronize: bool = True, + ) -> dict[str, Any]: + """Apply a FIFO batch of S3 sparse deltas through one collective RPC.""" + if self.llm is None: + raise RuntimeError( + "Attempting to update weights with either an uninitialized vLLM " + "or non-model-owner" + ) + if not serialized_payloads: + raise ValueError("A sparse refit batch must contain at least one payload.") + if not self._refit_workers_share_node: + results = [ + self._refit_collective_response( + self.llm.collective_rpc( + "update_weights_from_serialized_sparse_payload", + args=(payload, False), + ) + ) + for payload in serialized_payloads + ] + if synchronize: + self.llm.collective_rpc("synchronize_device", args=()) + timing: dict[str, float] = {} + merge_vllm_refit_receiver_timing(timing, results, maximum=False) + return {"ok": True, "payloads": len(serialized_payloads), **timing} + + paths: list[str] = [] + try: + for payload in serialized_payloads: + fd, path = tempfile.mkstemp( + prefix="nemo_rl_refit_", dir=self._refit_batch_staging_dir + ) + paths.append(path) + with os.fdopen(fd, "wb") as staged: + staged.write(payload) + response = self._refit_collective_response( + self.llm.collective_rpc( + "update_weights_from_sparse_payload_files", + args=tuple(paths), + kwargs={"synchronize": synchronize}, + ) + ) + finally: + for path in paths: + os.unlink(path) + response["payloads"] = len(serialized_payloads) + return response + def reset_prefix_cache(self): """Reset the prefix cache of vLLM engine.""" assert self.llm is not None, ( @@ -1090,6 +1361,15 @@ def wake_up(self, **kwargs): def shutdown(self) -> bool: """Clean up vLLM resources.""" try: + if self.refit_server is not None: + self.refit_server.should_exit = True + + self._flush_queued_sparse_payloads() + self._refit_apply_executor.shutdown(wait=True) + + if self.refit_server_thread is not None: + self.refit_server_thread.join(timeout=5.0) + if self.llm is not None: # Clean up extension resources (e.g., ZMQ sockets) self.llm.collective_rpc("cleanup", args=tuple()) diff --git a/nemo_rl/models/policy/lm_policy.py b/nemo_rl/models/policy/lm_policy.py index 397b4e086b5..c1f1a90084b 100644 --- a/nemo_rl/models/policy/lm_policy.py +++ b/nemo_rl/models/policy/lm_policy.py @@ -940,6 +940,12 @@ 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 init_remote_sparse_delta_baseline(self) -> list[ray.ObjectRef]: + """Initialize source-side sparse-delta baselines for remote S3 refit.""" + return self._run_s3_refit_workers( + "init_remote_sparse_delta_baseline", + ) + def finish_inference(self) -> None: """Offload policy model to CPU after inference.""" futures = self.worker_group.run_all_workers_single_data("finish_inference") @@ -1063,6 +1069,41 @@ def set_rollout_num_gpus_per_engine(self, num_gpus_per_engine: int) -> None: ) ) + def stream_sparse_weights_via_s3_manifest( + self, + refit_urls: list[str], + *, + api_key_env_var: Optional[str] = None, + timeout_s: float = 600.0, + ) -> list[ray.ObjectRef]: + """Upload vLLM refit payloads to S3 and post receiver manifests.""" + return self._run_s3_refit_workers( + "stream_sparse_weights_via_s3_manifest", + refit_urls=refit_urls, + api_key_env_var=api_key_env_var, + timeout_s=timeout_s, + ) + + def finish_remote_sparse_delta_sync(self, succeeded: bool) -> list[ray.ObjectRef]: + return self.worker_group.run_all_workers_single_data( + "finish_remote_sparse_delta_sync", succeeded=succeeded + ) + + def _run_s3_refit_workers( + self, + method_name: str, + **common_kwargs: Any, + ) -> list[ray.ObjectRef]: + worker_count = len(self.worker_group.workers) + return self.worker_group.run_all_workers_multiple_data( + method_name, + common_kwargs={ + **common_kwargs, + "shard_count": worker_count, + }, + shard_rank=list(range(worker_count)), + ) + def broadcast_weights_for_collective( self, kv_scales: Optional[dict[str, float]] = None ) -> list[ray.ObjectRef]: diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 6eea8bc4f8e..a1aaabd8af7 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -99,6 +99,11 @@ from nemo_rl.utils.packed_tensor import packed_broadcast_producer from nemo_rl.utils.r3_trace import maybe_r3_trace_stage from nemo_rl.utils.timer import Timer +from nemo_rl.utils.weight_transfer_s3_manifest import ( + init_sparse_delta_baseline_from_iterator, + stream_sparse_delta_payloads_via_s3_manifest, +) +from nemo_rl.utils.weight_transfer_sparse_codec import DeltaCompressionTracker TokenizerType = TypeVar("TokenizerType", bound=PreTrainedTokenizerBase) @@ -364,6 +369,12 @@ def __init__( self.is_generation_colocated = runtime_config.is_generation_colocated self.final_padded_vocab_size = runtime_config.final_padded_vocab_size self.sampling_params = runtime_config.sampling_params + delta_config = self.cfg.get("generation", {}).get("delta_compression") + self.delta_weight_transfer_tracker = ( + DeltaCompressionTracker(delta_config) + if delta_config and delta_config["enabled"] + else None + ) self.defer_fp32_logits = self.cfg["megatron_cfg"].get( "defer_fp32_logits", None @@ -1795,6 +1806,53 @@ def _get_model_config(self): return model.config return None + @torch.no_grad() + @wrap_with_nvtx_name("megatron_policy_worker/init_remote_sparse_delta_baseline") + def init_remote_sparse_delta_baseline( + self, + *, + shard_rank: int = 0, + shard_count: int = 1, + ) -> None: + """Initialize the source-side baseline for remote sparse S3 refit.""" + init_sparse_delta_baseline_from_iterator( + self._iter_params_with_optional_kv_scales(), + delta_tracker=self.delta_weight_transfer_tracker, + shard_rank=shard_rank, + shard_count=shard_count, + ) + + @torch.no_grad() + @wrap_with_nvtx_name("megatron_policy_worker/stream_sparse_weights_via_s3_manifest") + def stream_sparse_weights_via_s3_manifest( + self, + refit_urls: list[str], + *, + api_key_env_var: Optional[str] = None, + timeout_s: float = 600.0, + shard_rank: int = 0, + shard_count: int = 1, + ) -> dict[str, Any]: + """Upload vLLM refit payloads to S3 and post receiver manifests.""" + return stream_sparse_delta_payloads_via_s3_manifest( + self._iter_params_with_optional_kv_scales(), + delta_tracker=self.delta_weight_transfer_tracker, + refit_urls=refit_urls, + api_key_env_var=api_key_env_var, + timeout_s=timeout_s, + shard_rank=shard_rank, + shard_count=shard_count, + ) + + def finish_remote_sparse_delta_sync(self, *, succeeded: bool) -> None: + tracker = self.delta_weight_transfer_tracker + if tracker is None: + raise RuntimeError("Sparse delta tracker is not initialized.") + if succeeded: + tracker.on_sync_succeeded() + else: + tracker.on_sync_failed() + def _calculate_refit_param_info(self) -> list[tuple[str, int]]: """Calculate parameter information for refit. diff --git a/nemo_rl/utils/weight_transfer_s3.py b/nemo_rl/utils/weight_transfer_s3.py new file mode 100644 index 00000000000..1784e48ab0e --- /dev/null +++ b/nemo_rl/utils/weight_transfer_s3.py @@ -0,0 +1,122 @@ +# 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. + +"""AWS CRT object transport for S3 refit payloads.""" + +import io +from functools import cache +from typing import Any +from urllib.parse import quote + +_PART_SIZE = 64 * 1024**2 +_MEMORY_LIMIT = 2 * 1024**3 + + +@cache +def _s3_client(region: str) -> Any: + from awscrt.auth import AwsCredentialsProvider + from awscrt.io import ClientBootstrap, DefaultHostResolver, EventLoopGroup + from awscrt.s3 import S3Client, create_default_s3_signing_config + + event_loop_group = EventLoopGroup() + bootstrap = ClientBootstrap( + event_loop_group, + DefaultHostResolver(event_loop_group), + ) + credentials = AwsCredentialsProvider.new_default_chain(bootstrap) + return S3Client( + bootstrap=bootstrap, + region=region, + signing_config=create_default_s3_signing_config( + region=region, + credential_provider=credentials, + ), + part_size=_PART_SIZE, + multipart_upload_threshold=_PART_SIZE, + throughput_target_gbps=10.0, + memory_limit=_MEMORY_LIMIT, + ) + + +class S3ObjectStore: + """Blocking object operations backed by CRT's asynchronous S3 client.""" + + def __init__(self, *, bucket: str, region: str) -> None: + from awscrt.s3 import S3RequestType + + self.bucket = bucket + self.region = region + self._client = _s3_client(region) + self._request_type = S3RequestType + + def put_object(self, key: str, body: bytes) -> None: + request = self._request("PUT", key, body) + self._client.make_request( + type=self._request_type.PUT_OBJECT, + request=request, + ).finished_future.result() + + def get_object(self, key: str) -> bytes: + from awscrt.http import HttpHeaders + + body = bytearray() + + def on_headers( + status_code: int, + headers: list[tuple[str, str]], + **_kwargs: Any, + ) -> None: + nonlocal body + if status_code != 200: + raise RuntimeError(f"S3 GET returned HTTP {status_code}.") + length = HttpHeaders(headers).get("content-length") + if length is None: + raise RuntimeError("S3 GET response omitted content-length.") + body = bytearray(int(length)) + + def on_body(chunk: bytes, offset: int, **_kwargs: Any) -> None: + body[offset : offset + len(chunk)] = chunk + + self._client.make_request( + type=self._request_type.GET_OBJECT, + request=self._request("GET", key), + on_headers=on_headers, + on_body=on_body, + ).finished_future.result() + return bytes(body) + + def delete_object(self, key: str) -> None: + self._client.make_request( + type=self._request_type.DEFAULT, + request=self._request("DELETE", key), + operation_name="DeleteObject", + ).finished_future.result() + + def _request(self, method: str, key: str, body: bytes | None = None) -> Any: + from awscrt.http import HttpHeaders, HttpRequest + + headers = HttpHeaders( + [("host", f"{self.bucket}.s3.{self.region}.amazonaws.com")] + ) + if body is not None: + headers.add("content-length", str(len(body))) + headers.add("content-type", "application/octet-stream") + elif method == "DELETE": + headers.add("content-length", "0") + return HttpRequest( + method, + f"/{quote(key, safe='/~')}", + headers, + io.BytesIO(body) if body is not None else None, + ) diff --git a/nemo_rl/utils/weight_transfer_s3_manifest.py b/nemo_rl/utils/weight_transfer_s3_manifest.py new file mode 100644 index 00000000000..3ad35aa058f --- /dev/null +++ b/nemo_rl/utils/weight_transfer_s3_manifest.py @@ -0,0 +1,430 @@ +# 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. + +"""S3 manifest control-plane helpers for sparse vLLM refit.""" + +import io +import os +import threading +import time +import uuid +from collections.abc import Iterable, Iterator, Mapping, Sequence +from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait +from contextlib import suppress +from functools import cache +from typing import Any + +import requests +import torch +import zstandard +from urllib3.util.retry import Retry + +from nemo_rl.utils.packed_tensor import get_target_packed_tensor_size +from nemo_rl.utils.weight_transfer_s3 import S3ObjectStore +from nemo_rl.utils.weight_transfer_sparse_codec import ( + DeltaCompressionTracker, + NamedTensor, + TensorBatch, +) + +G_VLLM_REFIT_S3_MANIFEST_PATH = "/nemo-rl/refit/s3-manifest" +G_VLLM_REFIT_FLUSH_PATH = "/nemo-rl/refit/flush" +G_VLLM_REFIT_API_KEY_HEADER = "x-nemo-rl-refit-key" +_CONTROL_SESSION_LOCAL = threading.local() + + +def _env_int(name: str, default: int, min_value: int = 1) -> int: + value = int(os.getenv(name) or default) + if value < min_value: + raise ValueError(f"{name} must be >= {min_value}.") + return value + + +def _iter_chunks( + tensors: Iterable[NamedTensor], target_bytes: int +) -> Iterator[tuple[TensorBatch, float]]: + iterator = iter(tensors) + pending = None + while True: + started = time.perf_counter() + chunk = [pending] if pending is not None else [] + size = pending[1].numel() * pending[1].element_size() if pending else 0 + pending = None + for item in iterator: + item_size = item[1].numel() * item[1].element_size() + if chunk and size + item_size > target_bytes: + pending = item + break + chunk.append(item) + size += item_size + if size >= target_bytes: + break + export_pull_s = time.perf_counter() - started + if not chunk: + return + yield chunk, export_pull_s + + +def _http_session() -> requests.Session: + session = getattr(_CONTROL_SESSION_LOCAL, "session", None) + if session is None: + session = requests.Session() + adapter = requests.adapters.HTTPAdapter( + pool_connections=64, + pool_maxsize=64, + max_retries=Retry( + total=3, + backoff_factor=0.25, + status_forcelist=(500, 502, 503, 504), + allowed_methods={"POST"}, + ), + ) + session.mount("http://", adapter) + session.mount("https://", adapter) + _CONTROL_SESSION_LOCAL.session = session + return session + + +@cache +def _get_manifest_s3_store(bucket: str, region: str) -> S3ObjectStore: + return S3ObjectStore(bucket=bucket, region=region) + + +def vllm_refit_api_key(api_key_env_var: str | None) -> str | None: + if not api_key_env_var: + return None + token = os.environ.get(api_key_env_var) + if not token: + raise RuntimeError( + "vLLM S3 refit API key env var " + f"{api_key_env_var!r} is configured but unset or empty." + ) + return token + + +def _s3_sparse_export_chunk_size(delta_tracker: DeltaCompressionTracker) -> int: + requested = _env_int( + "NRL_REFIT_S3_EXPORT_CHUNK_BYTES", + default=256 * 1024**2, + min_value=1, + ) + if torch.cuda.is_available(): + requested = min(requested, get_target_packed_tensor_size()) + return min(requested, delta_tracker.sparse_bucket_size_bytes) + + +@cache +def _executor(key: str, workers: int) -> ThreadPoolExecutor: + return ThreadPoolExecutor(max_workers=workers, thread_name_prefix=f"nrl-{key}") + + +def _require_delta_tracker( + delta_tracker: DeltaCompressionTracker | None, +) -> DeltaCompressionTracker: + if delta_tracker is None: + raise RuntimeError("vLLM S3 sparse refit requires delta compression.") + return delta_tracker + + +def init_sparse_delta_baseline_from_iterator( + iterator: Iterable[NamedTensor], + *, + delta_tracker: DeltaCompressionTracker | None, + shard_rank: int = 0, + shard_count: int = 1, +) -> None: + start_s = time.perf_counter() + delta_tracker = _require_delta_tracker(delta_tracker) + export_chunk_size = _s3_sparse_export_chunk_size(delta_tracker) + + chunk_count = 0 + export_pull_s = snapshot_s = 0.0 + for chunk_index, (chunk, pull_s) in enumerate( + _iter_chunks(iterator, export_chunk_size) + ): + chunk_count = chunk_index + 1 + export_pull_s += pull_s + if chunk_index % shard_count != shard_rank: + continue + started = time.perf_counter() + delta_tracker.snapshot_baseline(chunk) + snapshot_s += time.perf_counter() - started + print( + "REFIT_BASELINE_INIT " + f"event=end chunks={chunk_count} export_pull_s={export_pull_s:.3f} " + f"snapshot_s={snapshot_s:.3f} " + f"seconds={time.perf_counter() - start_s:.3f}", + flush=True, + ) + + +def stream_sparse_delta_payloads_via_s3_manifest( + iterator: Iterable[NamedTensor], + *, + delta_tracker: DeltaCompressionTracker | None, + refit_urls: Sequence[str], + api_key_env_var: str | None = None, + timeout_s: float = 600.0, + shard_rank: int = 0, + shard_count: int = 1, +) -> dict[str, Any]: + urls = [url.strip().rstrip("/") for url in refit_urls if url.strip()] + if not urls: + raise ValueError("At least one vLLM S3 refit URL is required.") + delta_tracker = _require_delta_tracker(delta_tracker) + + bucket = os.getenv("NRL_REFIT_S3_BUCKET", "").strip() + if not bucket: + raise RuntimeError("NRL_REFIT_S3_BUCKET must be set for S3 refit.") + store = _get_manifest_s3_store( + bucket, + os.getenv("NRL_REFIT_S3_REGION", "us-east-1").strip() or "us-east-1", + ) + endpoint_urls = [f"{url}{G_VLLM_REFIT_S3_MANIFEST_PATH}" for url in urls] + object_prefix = os.getenv("NRL_REFIT_S3_PREFIX", "nemo-rl-refit").strip("/") + run_prefix = ( + f"{object_prefix}/{uuid.uuid4().hex}" if object_prefix else uuid.uuid4().hex + ) + encode_workers = _env_int( + "NRL_REFIT_S3_ENCODE_WORKERS", + default=max(2, min(8, os.cpu_count() or 8)), + min_value=1, + ) + upload_workers = _env_int( + "NRL_REFIT_S3_UPLOAD_WORKERS", + default=max(4, min(32, os.cpu_count() or 32)), + min_value=1, + ) + pipeline_workers = max(encode_workers, upload_workers) + executor = _executor("refit-s3-pipeline", pipeline_workers) + encode_slots = threading.Semaphore(encode_workers) + export_chunk_size = _s3_sparse_export_chunk_size(delta_tracker) + + def process_chunk(chunk: TensorBatch, payload_index: int) -> dict[str, Any] | None: + with encode_slots: + started = time.perf_counter() + payload = delta_tracker.prepare_sparse_delta_payload(chunk) + encode_s = time.perf_counter() - started + if not payload[2]: + return None + started = time.perf_counter() + buffer = io.BytesIO() + torch.save(payload, buffer) + raw_body = buffer.getvalue() + serialize_s = time.perf_counter() - started + started = time.perf_counter() + body = _zstd_compress(raw_body) + compress_s = time.perf_counter() - started + + key = f"{run_prefix}/{payload_index:06d}.pt" + started = time.perf_counter() + store.put_object(key, body) + s3_put_s = time.perf_counter() - started + + manifest = { + "bucket": store.bucket, + "region": store.region, + "key": key, + } + try: + started = time.perf_counter() + responses = _post_refit_body_to_endpoint_urls( + endpoint_urls, + manifest, + api_key_env_var=api_key_env_var, + timeout_s=timeout_s, + ) + manifest_post_s = time.perf_counter() - started + finally: + with suppress(Exception): + store.delete_object(key) + + return { + "body_size": len(body), + "encode_s": encode_s, + "serialize_s": serialize_s, + "compress_s": compress_s, + "s3_put_s": s3_put_s, + "manifest_post_s": manifest_post_s, + "receiver": merge_vllm_refit_receiver_timing({}, responses, maximum=True), + } + + timing: dict[str, float] = {} + receiver_timing: dict[str, float] = {} + counts = {"payloads": 0, "uploaded_bytes": 0} + chunk_count = 0 + export_pull_s = 0.0 + inflight: set[Any] = set() + max_inflight = pipeline_workers * 2 + + def collect_completed(future: Any) -> None: + result = future.result() + if result is None: + return + counts["payloads"] += 1 + counts["uploaded_bytes"] += int(result["body_size"]) + for key, value in result.items(): + if key.endswith("_s"): + timing[key] = timing.get(key, 0.0) + float(value) + merge_vllm_refit_receiver_timing( + receiver_timing, [result["receiver"]], maximum=False + ) + + def drain_completed() -> None: + completed, _ = wait(inflight, return_when=FIRST_COMPLETED) + for future in completed: + inflight.remove(future) + collect_completed(future) + + payload_index = 0 + stream_start = time.perf_counter() + try: + for chunk_index, (chunk, pull_s) in enumerate( + _iter_chunks(iterator, export_chunk_size) + ): + chunk_count = chunk_index + 1 + export_pull_s += pull_s + if chunk_index % shard_count != shard_rank: + continue + if len(inflight) >= max_inflight: + drain_completed() + inflight.add(executor.submit(process_chunk, chunk, payload_index)) + payload_index += 1 + + while inflight: + drain_completed() + except Exception: + for future in inflight: + future.cancel() + wait(inflight) + with suppress(Exception): + flush_vllm_refit_urls( + urls, + api_key_env_var=api_key_env_var, + timeout_s=min(timeout_s, 60.0), + ) + raise + + timing = { + "total_s": time.perf_counter() - stream_start, + "export_pull_s": export_pull_s, + **timing, + "payloads": counts["payloads"], + "chunks": chunk_count, + "uploaded_mb": counts["uploaded_bytes"] / 1e6, + "pipeline_workers": pipeline_workers, + "encode_workers": encode_workers, + "export_chunk_mb": export_chunk_size / 1e6, + "shard_rank": shard_rank, + "shard_count": shard_count, + } + timing.update(receiver_timing) + print( + "REFIT_S3_TIMING " + + " ".join(f"{key}={value}" for key, value in timing.items()), + flush=True, + ) + return {"ok": True, "payloads": counts["payloads"]} + + +def _post_refit_body_to_endpoint_urls( + endpoint_urls: Sequence[str], + body: Mapping[str, str], + *, + api_key_env_var: str | None, + timeout_s: float, +) -> list[dict[str, Any]]: + headers = {} + if token := vllm_refit_api_key(api_key_env_var): + headers[G_VLLM_REFIT_API_KEY_HEADER] = token + + def post(url: str) -> dict[str, Any]: + response = _http_session().post( + url, + json=body, + headers=headers, + timeout=timeout_s, + ) + result: dict[str, Any] = response.json() if response.content else {} + if response.status_code >= 400 or result.get("ok") is not True: + raise RuntimeError(f"vLLM refit failed for {url}: {result}") + return result + + return list(_executor("refit-fanout", len(endpoint_urls)).map(post, endpoint_urls)) + + +def flush_vllm_refit_urls( + base_urls: Sequence[str], + *, + api_key_env_var: str | None, + timeout_s: float, +) -> None: + endpoint_urls = [ + f"{url}{G_VLLM_REFIT_FLUSH_PATH}" + for url in (url.strip().rstrip("/") for url in base_urls if url.strip()) + ] + _post_refit_body_to_endpoint_urls( + endpoint_urls, + {}, + api_key_env_var=api_key_env_var, + timeout_s=timeout_s, + ) + + +def download_s3_refit_payload( + manifest: Mapping[str, Any], +) -> bytes: + bucket, region, key = ( + str(manifest[field]) for field in ("bucket", "region", "key") + ) + return _zstd_decompress(_get_manifest_s3_store(bucket, region).get_object(key)) + + +def merge_vllm_refit_receiver_timing( + result: dict[str, Any], + timings: Iterable[Mapping[str, Any]], + *, + maximum: bool, +) -> dict[str, Any]: + for timing in timings: + for key, value in timing.items(): + if key.startswith("receiver_") and key.endswith("_s"): + number = float(value) + if key in result: + number = ( + max(float(result[key]), number) + if maximum + else float(result[key]) + number + ) + result[key] = number + return result + + +def _zstd_compress(raw: bytes) -> bytes: + compressor = getattr(_CONTROL_SESSION_LOCAL, "zstd_compressor", None) + if compressor is None: + compressor = zstandard.ZstdCompressor( + level=1, + threads=_env_int("NRL_REFIT_S3_ZSTD_THREADS", default=0, min_value=0), + ) + _CONTROL_SESSION_LOCAL.zstd_compressor = compressor + return compressor.compress(raw) + + +def _zstd_decompress(raw: bytes) -> bytes: + decompressor = getattr(_CONTROL_SESSION_LOCAL, "zstd_decompressor", None) + if decompressor is None: + decompressor = zstandard.ZstdDecompressor() + _CONTROL_SESSION_LOCAL.zstd_decompressor = decompressor + return decompressor.decompress(raw) diff --git a/nemo_rl/utils/weight_transfer_sparse_codec.py b/nemo_rl/utils/weight_transfer_sparse_codec.py new file mode 100644 index 00000000000..a0cafc5d8c4 --- /dev/null +++ b/nemo_rl/utils/weight_transfer_sparse_codec.py @@ -0,0 +1,245 @@ +# 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 os +import tempfile +import threading +from collections.abc import Iterable, Mapping +from concurrent.futures import ThreadPoolExecutor +from typing import Any + +import numpy as np +import torch + +NamedTensor = tuple[str, torch.Tensor] +TensorBatch = list[NamedTensor] +TensorPayload = tuple[torch.Tensor, torch.Tensor, list[dict[str, Any]]] + + +def encode_sparse_infos( + infos: Iterable[tuple[str, torch.Tensor, torch.Tensor, torch.Tensor]], + *, + empty_dtype: torch.dtype, +) -> TensorPayload: + packed_locations = [] + packed_values = [] + metadata: list[dict[str, Any]] = [] + index_offset = value_offset = 0 + for name, tensor, raw_locations, raw_values in infos: + count = int(raw_values.numel()) + if count == 1 or int(raw_locations[-1] - raw_locations[0] + 1) == count: + index_count = 0 + location_metadata = { + "index_encoding": "range", + "range_start": int(raw_locations[0]), + } + else: + location_tensor = _encode_explicit_locations(raw_locations) + location_metadata = {"index_encoding": "deltas"} + packed_locations.append(location_tensor) + index_count = int(location_tensor.numel()) + packed_values.append(raw_values) + metadata.append( + { + "name": name, + "shape": tuple(int(dim) for dim in tensor.shape), + "index_start": index_offset, + "index_end": index_offset + index_count, + "value_start": value_offset, + "value_end": value_offset + count, + **location_metadata, + } + ) + index_offset += index_count + value_offset += count + indices = ( + torch.cat(packed_locations) + if packed_locations + else torch.empty(0, dtype=torch.uint8) + ) + values = ( + torch.cat(packed_values) if packed_values else torch.empty(0, dtype=empty_dtype) + ) + return indices, values, metadata + + +def sparse_locations_for_item( + item: dict[str, Any], + packed_locations: torch.Tensor, + *, + device: torch.device | int | str, +) -> torch.Tensor: + count = int(item["value_end"]) - int(item["value_start"]) + if item["index_encoding"] == "range": + start = int(item["range_start"]) + return torch.arange(start, start + count, device=device) + + index_start, index_end = int(item["index_start"]), int(item["index_end"]) + raw = ( + packed_locations[index_start:index_end] + .detach() + .cpu() + .numpy() + .astype(np.uint8, copy=False) + .tobytes() + ) + delta_dtype = {2: np.uint16, 4: np.uint32, 8: np.uint64}[len(raw) // count] + deltas = np.frombuffer(raw, dtype=delta_dtype).astype(np.int64, copy=False) + locations = np.cumsum(deltas + 1, dtype=np.int64) - 1 + return torch.from_numpy(locations).to(device=device) + + +def _encode_explicit_locations( + locations: torch.Tensor, +) -> torch.Tensor: + indices = locations.detach().cpu().numpy().astype(np.int64, copy=False) + deltas = np.diff(indices, prepend=-1) - 1 + max_delta = int(deltas.max()) + dtype = next( + dtype + for dtype in (np.uint16, np.uint32, np.uint64) + if max_delta <= np.iinfo(dtype).max + ) + raw = deltas.astype(dtype, copy=False).tobytes() + return torch.from_numpy(np.frombuffer(raw, dtype=np.uint8).copy()) + + +class DeltaCompressionTracker: + """Source-side CPU or mmap baseline for sparse-delta refit.""" + + def __init__(self, config: Mapping[str, Any]) -> None: + self.sparse_bucket_size_bytes = int(config["sparse_bucket_size_bytes"]) + if self.sparse_bucket_size_bytes < 1: + raise ValueError("delta_compression.sparse_bucket_size_bytes must be >= 1") + dtype_name = {"bf16": "bfloat16", "fp16": "float16", "fp32": "float32"}.get( + str(config["dtype"]).lower(), str(config["dtype"]).lower() + ) + self.delta_dtype: torch.dtype = getattr(torch, dtype_name) + self.baseline_in_memory = os.getenv("NRL_REFIT_BASELINE_IN_MEMORY") == "1" + self.baseline_mmap_dir = config.get("baseline_mmap_dir") or os.getenv( + "NRL_REFIT_BASELINE_MMAP_DIR" + ) + self.baseline: dict[str, torch.Tensor] = {} + self._pending_updates: dict[str, tuple[torch.Tensor, torch.Tensor]] = {} + self._pending_updates_lock = threading.Lock() + self._baseline_commits: tuple[Any, ...] = () + self._baseline_commit_lock = threading.Lock() + self._baseline_commit_executor = ThreadPoolExecutor( + max_workers=4, thread_name_prefix="nrl-refit-baseline" + ) + + def prepare_sparse_delta_payload(self, tensors: TensorBatch) -> TensorPayload: + self._wait_for_baseline_commits() + sparse_infos = [] + pending_updates = {} + for name, tensor in tensors: + baseline = self.baseline.get(name) + if baseline is None: + raise RuntimeError(f"Sparse delta baseline is missing {name!r}.") + current = tensor.detach().cpu() + current_flat, baseline_flat = current.view(-1), baseline.view(-1) + locations = (current_flat != baseline_flat).nonzero().view(-1) + if locations.numel(): + current_values = current_flat[locations] + sparse_infos.append( + ( + name, + current, + locations, + (current_values - baseline_flat[locations]).to( + self.delta_dtype + ), + ) + ) + pending_updates[name] = (locations, current_values) + with self._pending_updates_lock: + self._pending_updates.update(pending_updates) + return encode_sparse_infos(sparse_infos, empty_dtype=self.delta_dtype) + + def on_sync_succeeded(self) -> None: + with self._pending_updates_lock: + pending_updates, self._pending_updates = self._pending_updates, {} + items = list(pending_updates.items()) + workers = min(4, len(items)) + with self._baseline_commit_lock: + self._baseline_commits = tuple( + self._baseline_commit_executor.submit( + self._commit_baseline_updates, items[worker::workers] + ) + for worker in range(workers) + ) + + def on_sync_failed(self) -> None: + with self._pending_updates_lock: + self._pending_updates.clear() + + def snapshot_baseline(self, tensors: Iterable[NamedTensor]) -> None: + self._wait_for_baseline_commits() + for name, tensor in tensors: + self._baseline(name, tuple(tensor.shape), tensor.dtype).copy_(tensor) + + def _wait_for_baseline_commits(self) -> None: + with self._baseline_commit_lock: + commits = self._baseline_commits + for commit in commits: + commit.result() + with self._baseline_commit_lock: + if self._baseline_commits == commits: + self._baseline_commits = () + + def _commit_baseline_updates( + self, updates: Iterable[tuple[str, tuple[torch.Tensor, torch.Tensor]]] + ) -> None: + for name, (locations, values) in updates: + target = self.baseline[name].view(-1) + count = locations.numel() + if count > 1: + first, last = int(locations[0]), int(locations[-1]) + span = last - first + if span % (count - 1) == 0: + step = span // (count - 1) + if step == 1 or all( + torch.equal( + locations[start:end], + first + + torch.arange(start, end, dtype=locations.dtype) * step, + ) + for start in range(0, count, 1 << 20) + for end in (min(start + (1 << 20), count),) + ): + target[first : last + 1 : step].copy_(values) + continue + target.index_copy_(0, locations, values) + + def _baseline( + self, + name: str, + shape: tuple[int, ...], + dtype: torch.dtype, + ) -> torch.Tensor: + if name in self.baseline: + return self.baseline[name] + if self.baseline_in_memory: + baseline = torch.empty(shape, dtype=dtype) + else: + numel = torch.Size(shape).numel() + with tempfile.NamedTemporaryFile( + prefix="nrl-refit-baseline-", dir=self.baseline_mmap_dir + ) as handle: + handle.truncate(numel * torch.empty((), dtype=dtype).element_size()) + baseline = torch.from_file( + handle.name, shared=True, size=numel, dtype=dtype + ).view(shape) + self.baseline[name] = baseline + return baseline diff --git a/nemo_rl/weight_sync/vllm_s3_sparse_weight_synchronizer.py b/nemo_rl/weight_sync/vllm_s3_sparse_weight_synchronizer.py new file mode 100644 index 00000000000..3c3c36109db --- /dev/null +++ b/nemo_rl/weight_sync/vllm_s3_sparse_weight_synchronizer.py @@ -0,0 +1,122 @@ +# 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. + +"""S3 manifest weight synchronizer for remote non-colocated vLLM refit.""" + +import time +from contextlib import nullcontext +from typing import Any + +import ray + +from nemo_rl.utils.timer import Timer +from nemo_rl.utils.weight_transfer_s3_manifest import flush_vllm_refit_urls +from nemo_rl.weight_sync.interfaces import WeightSynchronizer + + +class VllmS3SparseWeightSynchronizer(WeightSynchronizer): + def __init__( + self, + policy: Any, + generation: Any, + *, + api_key_env_var: str | None = None, + request_timeout_s: float = 600.0, + ) -> None: + self._policy = policy + self._generation = generation + self._refit_urls: list[str] = [] + self._api_key_env_var = api_key_env_var + self._request_timeout_s = request_timeout_s + self._stale = True + self._baseline_init_refs: list[Any] | None = None + self._baseline_commit_refs: list[Any] | None = None + + def sync_weights( + self, + *, + timer: Timer | None = None, + kv_scales: dict[str, float] | None = None, + ) -> None: + timer_context = ( + timer.time("prepare_for_generation/transfer_and_update_weights") + if timer is not None + else nullcontext() + ) + with timer_context: + if self._baseline_commit_refs is not None: + ray.get(self._baseline_commit_refs) + self._baseline_commit_refs = None + flush_success = self._generation.invalidate_kv_cache() + if not flush_success: + print("vLLM KV cache invalidation failed before S3 weight update.") + + if self._baseline_init_refs is not None: + ray.get(self._baseline_init_refs) + self._baseline_init_refs = None + succeeded = False + try: + results = ray.get( + self._policy.stream_sparse_weights_via_s3_manifest( + self._refit_urls, + api_key_env_var=self._api_key_env_var, + timeout_s=self._request_timeout_s, + ) + ) + payloads = sum(int(result["payloads"]) for result in results) + if payloads: + started = time.perf_counter() + flush_vllm_refit_urls( + self._refit_urls, + api_key_env_var=self._api_key_env_var, + timeout_s=self._request_timeout_s, + ) + print( + "REFIT_S3_GLOBAL_FLUSH " + f"payloads={payloads} " + f"seconds={time.perf_counter() - started:.3f}", + flush=True, + ) + succeeded = True + finally: + self._baseline_commit_refs = ( + self._policy.finish_remote_sparse_delta_sync(succeeded) + ) + self._stale = False + + @property + def is_stale(self) -> bool: + return self._stale + + def mark_stale(self) -> None: + self._stale = True + + def init_communicator(self) -> None: + self._baseline_init_refs = self._policy.init_remote_sparse_delta_baseline() + self._refit_urls = self._generation.report_refit_server_base_urls() + if not self._refit_urls: + raise ValueError( + "vLLM S3 sparse refit requires expose_http_refit_server=true." + ) + self._stale = False + + def shutdown(self) -> None: + for ref in (self._baseline_init_refs or []) + ( + self._baseline_commit_refs or [] + ): + ray.cancel(ref, force=False) + self._baseline_init_refs = None + self._baseline_commit_refs = None + self._refit_urls = [] + self._stale = True diff --git a/pyproject.toml b/pyproject.toml index 11dc7691676..63b1de062e3 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -66,6 +66,8 @@ dependencies = [ "nccl4py; sys_platform != 'darwin'", # for non-colocated refit "cuda-bindings; sys_platform != 'darwin'", # for non-colocated refit "pybase64", # for sglang refit + "awscrt>=0.35.0", # for parallel S3 refit transport + "zstandard", # for sparse refit body compression "nvidia-cudnn-cu13==9.20.0.48; sys_platform != 'darwin'", # for transformer-engine no build isolation # tilelang — replacement Triton kernel mamba-ssm requires when # Triton >= 3.4.0 on Hopper, see github.com/state-spaces/mamba#640. diff --git a/pyrefly.toml b/pyrefly.toml index 117ec1fdcd7..2a304c21672 100644 --- a/pyrefly.toml +++ b/pyrefly.toml @@ -18,6 +18,7 @@ replace-imports-with-any = [ "numpy.*", "sphinx.*", "docutils.*", + "zstandard.*", ] project-includes = [ # TODO: enable these once we have 100 correctness @@ -199,6 +200,10 @@ project-includes = [ "nemo_rl/weight_sync/http_weight_synchronizer.py", "nemo_rl/weight_sync/interfaces.py", "nemo_rl/weight_sync/ipc_weight_synchronizer.py", + "nemo_rl/weight_sync/vllm_s3_sparse_weight_synchronizer.py", + "nemo_rl/utils/weight_transfer_s3_manifest.py", + "nemo_rl/utils/weight_transfer_s3.py", + "nemo_rl/utils/weight_transfer_sparse_codec.py", "tools/model_diagnostics/1.max_model_len_respected.py", "tools/model_diagnostics/2.long_generation_decode_vs_prefill.py", "tools/model_diagnostics/3.check_and_reinit_hf_model_embeddings_untrained.py", diff --git a/tests/unit/models/generation/test_vllm_backend.py b/tests/unit/models/generation/test_vllm_backend.py index 2fd83278116..a63242a9f0f 100644 --- a/tests/unit/models/generation/test_vllm_backend.py +++ b/tests/unit/models/generation/test_vllm_backend.py @@ -19,6 +19,7 @@ import contextlib import json from types import SimpleNamespace +from typing import Any from unittest.mock import MagicMock import pytest @@ -87,6 +88,237 @@ def _patch_vllm_postload(monkeypatch): return process_weights +def _attach_tensor_attrs(tensor: torch.Tensor, **attrs: object) -> torch.Tensor: + for name, value in attrs.items(): + setattr(tensor, name, value) + return tensor + + +def _make_sparse_delta_extension( + parameter_name: str, + target: torch.Tensor, + module: object, + module_name: str | None = None, +) -> Any: + from nemo_rl.models.generation.vllm.vllm_backend import ( + VllmInternalWorkerExtension, + ) + + ext = VllmInternalWorkerExtension.__new__(VllmInternalWorkerExtension) + ext.rank = 1 + ext._direct_sparse_delta_targets = {parameter_name: target} + ext._direct_sparse_delta_modules = { + module_name or parameter_name.rsplit(".", 1)[0]: module + } + ext._direct_sparse_delta_plan_cache = {} + return ext + + +def _assert_sparse_plan( + ext: Any, + plan: Any, + source_locations: list[int], + expected_locations: list[int], + expected_values: list[float], +) -> None: + assert plan is not None + values = torch.arange(len(source_locations), dtype=torch.float32) + locations, values = ext._local_sparse_delta_update_inputs( + torch.tensor(source_locations), values, plan + ) + assert locations.tolist() == expected_locations + assert values.tolist() == expected_values + + +@pytest.mark.vllm +def test_serialized_sparse_payload_batch_preserves_order(tmp_path) -> None: + from nemo_rl.models.generation.vllm.vllm_backend import ( + VllmInternalWorkerExtension, + ) + + ext = VllmInternalWorkerExtension.__new__(VllmInternalWorkerExtension) + ext.device = torch.device("cpu") + payloads = [ + (torch.tensor([index]), torch.tensor([float(index)]), {"index": index}) + for index in range(3) + ] + paths = [tmp_path / f"{index}.pt" for index in range(3)] + for path, payload in zip(paths, payloads, strict=True): + torch.save(payload, path) + applied: list[tuple[Any, bool]] = [] + + def apply(payload: Any, *, synchronize: bool) -> dict[str, Any]: + applied.append((payload, synchronize)) + return { + "ok": True, + "receiver_sparse_apply_s": 2.0, + } + + ext._apply_sparse_request = apply + result = ext.update_weights_from_sparse_payload_files( + *(str(path) for path in paths), synchronize=False + ) + + assert [item[0][2]["index"] for item in applied] == [0, 1, 2] + assert all( + torch.equal(item[0][1], payload[1]) + for item, payload in zip(applied, payloads, strict=True) + ) + assert all(not synchronize for _, synchronize in applied) + assert result["receiver_deserialize_s"] >= 0.0 + assert result["receiver_sparse_apply_s"] == 6.0 + + +@pytest.mark.vllm +def test_direct_sparse_delta_placement() -> None: + qkv_name = "model.layers.0.self_attn.qkv_proj.weight" + qkv_target = _attach_tensor_attrs(torch.zeros(8, 2), output_dim=0) + ext = _make_sparse_delta_extension( + qkv_name, + qkv_target, + SimpleNamespace( + tp_rank=1, + num_kv_head_replicas=2, + _get_shard_offset_mapping=lambda shard: {"q": 0, "k": 4, "v": 6}[shard], + _get_shard_size_mapping=lambda shard: {"q": 4, "k": 2, "v": 2}[shard], + ), + ) + qkv_source = "model.layers.0.self_attn.k_proj.weight" + plan = ext._direct_sparse_delta_qkv_plan( + {"name": qkv_source, "shape": (2, 2)}, qkv_source, {qkv_name: qkv_target} + ) + _assert_sparse_plan(ext, plan, [0, 1, 2, 3], [8, 9, 10, 11], [0.0, 1.0, 2.0, 3.0]) + + expert_name = "model.layers.0.mlp.experts.w13_weight" + expert_target = torch.zeros(2, 4, 2) + expert_module = SimpleNamespace( + tp_rank=1, + moe_config=SimpleNamespace(is_act_and_mul=False), + _map_global_expert_id_to_local_expert_id=lambda expert: ( + 1 if expert == 3 else -1 + ), + ) + ext = _make_sparse_delta_extension( + expert_name, + expert_target, + expert_module, + ) + expert_source = "model.layers.0.mlp.experts.3.gate_proj.weight" + plan = ext._direct_sparse_delta_expert_plan( + {"name": expert_source, "shape": (8, 2)}, + expert_source, + {expert_name: expert_target}, + ) + _assert_sparse_plan( + ext, + plan, + [6, 7, 8, 9, 14, 15], + [8, 9, 14, 15], + [2.0, 3.0, 4.0, 5.0], + ) + + expert_source = "model.layers.0.mlp.experts.3.up_proj.weight" + plan = ext._direct_sparse_delta_expert_plan( + {"name": expert_source, "shape": (8, 2)}, + expert_source, + {expert_name: expert_target}, + ) + _assert_sparse_plan( + ext, + plan, + [6, 7, 8, 9, 14, 15], + [8, 9, 14, 15], + [2.0, 3.0, 4.0, 5.0], + ) + + w2_target = torch.zeros(2, 2, 4) + ext = _make_sparse_delta_extension(expert_name, w2_target, expert_module) + expert_source = "model.layers.0.mlp.experts.3.down_proj.weight" + plan = ext._direct_sparse_delta_expert_plan( + {"name": expert_source, "shape": (2, 8)}, + expert_source, + {"model.layers.0.mlp.experts.w2_weight": w2_target}, + ) + _assert_sparse_plan(ext, plan, [3, 4, 7, 11, 15], [8, 11, 15], [1.0, 2.0, 4.0]) + + mamba_name = "model.layers.0.mixer.in_proj.weight" + mamba_target = torch.zeros(16, 1, 2) + ext = _make_sparse_delta_extension( + mamba_name, + mamba_target, + SimpleNamespace( + tp_size=2, + intermediate_size=8, + groups_ssm_state_size=6, + num_heads=4, + ), + "model.layers.0.mixer", + ) + plan = ext._direct_sparse_delta_mamba2_plan( + {"name": mamba_name, "shape": (28, 2)}, + mamba_name, + {mamba_name: mamba_target}, + ) + _assert_sparse_plan( + ext, plan, [0, 8, 25, 38, 40, 55], [0, 9, 22, 31], [1.0, 2.0, 4.0, 5.0] + ) + + mamba_target = torch.zeros(14, 2) + ext = _make_sparse_delta_extension( + mamba_name, + mamba_target, + SimpleNamespace( + tp_size=2, + intermediate_size=8, + groups_ssm_state_size=4, + num_heads=4, + ), + "model.layers.0.mixer", + ) + plan = ext._direct_sparse_delta_mamba2_plan( + {"name": mamba_name, "shape": (28, 2)}, + mamba_name, + {mamba_name: mamba_target}, + ) + _assert_sparse_plan( + ext, + plan, + [0, 8, 24, 36, 44, 52, 55], + [0, 8, 16, 20, 24, 27], + [1, 2, 3, 4, 5, 6], + ) + + shard_target = _attach_tensor_attrs( + torch.zeros(3, 2), output_dim=0, tp_size=2, tp_rank=1 + ) + ext = _make_sparse_delta_extension("down_proj.weight", shard_target, object()) + plan = ext._direct_sparse_delta_shard_plan( + {"name": "down_proj.weight", "shape": (6, 2)}, shard_target + ) + _assert_sparse_plan( + ext, + plan, + [0, 1, 6, 7, 10, 11], + [0, 1, 4, 5], + [2.0, 3.0, 4.0, 5.0], + ) + + shared_target = _attach_tensor_attrs( + torch.zeros(3, 2), output_dim=0, input_dim=1, tp_size=2, tp_rank=1 + ) + ext = _make_sparse_delta_extension("down_proj.weight", shared_target, object()) + plan = ext._direct_sparse_delta_shard_plan( + {"name": "down_proj.weight", "shape": (3, 4)}, shared_target + ) + _assert_sparse_plan( + ext, + plan, + [0, 1, 2, 3, 6, 7, 10, 11], + [0, 1, 2, 3, 4, 5], + [2.0, 3.0, 4.0, 5.0, 6.0, 7.0], + ) + + @pytest.mark.vllm def test_update_weights_from_collective_processes_weights_after_loading(monkeypatch): from nemo_rl.models.generation.vllm import vllm_backend diff --git a/tests/unit/models/generation/test_vllm_generation.py b/tests/unit/models/generation/test_vllm_generation.py index ec659bf7602..de870b6325d 100644 --- a/tests/unit/models/generation/test_vllm_generation.py +++ b/tests/unit/models/generation/test_vllm_generation.py @@ -16,9 +16,13 @@ import json import os import sys +import threading import types +from collections import deque +from concurrent.futures import ThreadPoolExecutor from copy import deepcopy from pathlib import Path +from typing import Any from unittest.mock import MagicMock import pytest @@ -37,6 +41,8 @@ ) from nemo_rl.models.generation.vllm import VllmConfig, VllmGeneration from nemo_rl.models.generation.vllm.vllm_worker import ( + BaseVllmGenerationWorker, + VllmGenerationWorkerImpl, _resolve_enable_prefix_caching, ) from nemo_rl.models.generation.vllm.vllm_worker_async import ( @@ -159,6 +165,97 @@ def test_resolve_enable_prefix_caching_uses_cuda_capability_for_auto(monkeypatch assert _resolve_enable_prefix_caching({}) is False +def test_sparse_refit_queue_batches_payloads_in_fifo_order() -> None: + worker = BaseVllmGenerationWorker.__new__(BaseVllmGenerationWorker) + worker._refit_apply_queue_lock = threading.Lock() + worker._refit_apply_executor = ThreadPoolExecutor(max_workers=1) + worker._refit_apply_futures = deque() + worker._refit_apply_pending_payloads = [] + worker._refit_apply_payload_count = 0 + worker._refit_apply_batch_count = 0 + worker._refit_apply_queue_depth = 2 + worker._refit_apply_batch_size = 3 + worker.llm = MagicMock() + applied: list[tuple[tuple[bytes, ...], bool]] = [] + + def apply(payloads: tuple[bytes, ...], synchronize: bool) -> dict[str, Any]: + applied.append((payloads, synchronize)) + return { + "ok": True, + "payloads": len(payloads), + "receiver_total_s": float(len(payloads)), + } + + worker.update_weights_from_serialized_sparse_payloads = apply + try: + responses = [ + worker._enqueue_sparse_payload_apply(payload) + for payload in (b"0", b"1", b"2", b"3", b"4") + ] + response = worker._flush_queued_sparse_payloads() + responses.append(response) + finally: + worker._refit_apply_executor.shutdown(wait=True) + + assert applied == [ + ((b"0", b"1", b"2"), False), + ((b"3", b"4"), False), + ] + assert response["payloads"] == 5 + assert response["batches"] == 2 + assert sum(result.get("receiver_total_s", 0.0) for result in responses) == 5.0 + worker.llm.collective_rpc.assert_called_once_with("synchronize_device", args=()) + + +def test_sparse_refit_batch_uses_one_collective_rpc(tmp_path: Path) -> None: + worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) + staged_payloads: list[bytes] = [] + + def collective_rpc(method, args, kwargs): + assert method == "update_weights_from_sparse_payload_files" + staged_payloads.extend(Path(path).read_bytes() for path in args) + assert kwargs == {"synchronize": False} + return [{"ok": True, "receiver_total_s": 1.0}] + + worker.llm = MagicMock(collective_rpc=MagicMock(side_effect=collective_rpc)) + worker._refit_workers_share_node = True + worker._refit_batch_staging_dir = str(tmp_path) + payloads = (b"0", b"1", b"2") + + response = worker.update_weights_from_serialized_sparse_payloads( + payloads, synchronize=False + ) + + assert staged_payloads == list(payloads) + assert not list(tmp_path.iterdir()) + worker.llm.collective_rpc.assert_called_once() + assert response == {"ok": True, "receiver_total_s": 1.0, "payloads": 3} + + +def test_sparse_refit_batch_falls_back_across_nodes() -> None: + worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) + worker._refit_workers_share_node = False + worker.llm = MagicMock( + collective_rpc=MagicMock(return_value=[{"ok": True, "receiver_total_s": 1.0}]) + ) + + response = worker.update_weights_from_serialized_sparse_payloads((b"0", b"1", b"2")) + + calls = worker.llm.collective_rpc.call_args_list + assert [call.args[0] for call in calls] == [ + "update_weights_from_serialized_sparse_payload", + "update_weights_from_serialized_sparse_payload", + "update_weights_from_serialized_sparse_payload", + "synchronize_device", + ] + assert [call.kwargs["args"] for call in calls[:3]] == [ + (b"0", False), + (b"1", False), + (b"2", False), + ] + assert response == {"ok": True, "receiver_total_s": 3.0, "payloads": 3} + + basic_lora_test_config: LoRAConfig = { "enabled": False, "target_modules": [], @@ -447,6 +544,17 @@ def test_configure_generation_config_uses_real_startup_weights_without_draft_ref assert configured["vllm_cfg"]["load_format"] == "auto" +def test_configure_generation_config_uses_real_s3_delta_baseline(): + vllm_config = deepcopy(basic_vllm_test_config) + vllm_config["refit_transport"] = "vllm_s3_sparse" + + configured = configure_generation_config( + vllm_config, MagicMock(pad_token_id=0, eos_token_id=1) + ) + + assert configured["vllm_cfg"]["load_format"] == "auto" + + def test_configure_generation_config_keeps_dummy_startup_weights_with_draft_refit(): """Speculative training can keep dummy startup weights when draft refit is available.""" vllm_config = deepcopy(basic_vllm_test_config) diff --git a/tests/unit/reference_configs/grpo_math_1B.yaml b/tests/unit/reference_configs/grpo_math_1B.yaml index 55cdc1d4247..48861f83257 100644 --- a/tests/unit/reference_configs/grpo_math_1B.yaml +++ b/tests/unit/reference_configs/grpo_math_1B.yaml @@ -332,6 +332,8 @@ policy: top_k: null stop_token_ids: null stop_strings: null + refit_transport: null # Set to "vllm_s3_sparse" to use S3 sparse-delta refit. + delta_compression: null # S3 sparse-delta refit config; null uses the existing refit path. mcore_generation_config: async_engine: false max_model_len: ${policy.max_total_sequence_length} @@ -369,6 +371,9 @@ policy: num_first_layers_in_bf16: 0 enable_vllm_metrics_logger: true # Set to true to enable vLLM internal metrics logger, turn off for better performance vllm_metrics_logger_interval: 0.5 # Interval in seconds to collect vLLM logger metrics + expose_http_refit_server: false # Start the internal sparse-delta refit endpoint on vLLM workers. + http_refit_api_key_env_var: null # Optional env var containing the internal refit API key. + http_refit_server_port: null # Optional fixed port for Kubernetes targetPorts. vllm_kwargs: {} colocated: # true: generation shares training GPUs diff --git a/uv.lock b/uv.lock index 1e5f744404a..0bfea3edf89 100644 --- a/uv.lock +++ b/uv.lock @@ -422,6 +422,24 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/64/b4/17d4b0b2a2dc85a6df63d1157e028ed19f90d4cd97c36717afef2bc2f395/attrs-26.1.0-py3-none-any.whl", hash = "sha256:c647aa4a12dfbad9333ca4e71fe62ddc36f4e63b2d260a37a8b83d2f043ac309", size = 67548, upload-time = "2026-03-19T14:22:23.645Z" }, ] +[[package]] +name = "awscrt" +version = "0.35.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/e8/8a/294c2f6cdda8f386057a5f6b349fec9f4838b9c25a98cb67dc503bb80514/awscrt-0.35.0.tar.gz", hash = "sha256:761ae0dda17fd9dfaff4bbb2a376e28e44dfd77dc6410b7bc408297a1fd5600e", size = 37016406, upload-time = "2026-06-25T18:17:26.611Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/d9/5a/44aa794eee204002ae161ddcd1a0d901c6c9ca587f22294e345bb468b3ae/awscrt-0.35.0-cp311-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:7cd20c94a6008164c89eeb420892c3d88b165a79bdd22ef0a4ee383bee8b4cdb", size = 3964249, upload-time = "2026-06-25T18:16:30.389Z" }, + { url = "https://files.pythonhosted.org/packages/84/03/edba2d4e7bf381eece2ff176836dbdbb5111d5a81f06705415a899848da9/awscrt-0.35.0-cp311-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:e26cca47c8b84dd968bde8c39385010acd710631e55b686a0e5422573c89a25b", size = 4258920, upload-time = "2026-06-25T18:16:31.611Z" }, + { url = "https://files.pythonhosted.org/packages/17/56/67c6374c97da326ef4a1af075b6a776cbcdc1bf14ddd259aa9acc613adfa/awscrt-0.35.0-cp311-abi3-musllinux_1_1_aarch64.whl", hash = "sha256:ab050c01fb3a64c4efc7a12baf49321c10e1635903e8ad05d1fe7a88ef5b0f2a", size = 3876609, upload-time = "2026-06-25T18:16:33.032Z" }, + { url = "https://files.pythonhosted.org/packages/2d/48/40014ff6278699e0065165128742a35b0654cb544e980496bd09b5d9c5db/awscrt-0.35.0-cp311-abi3-musllinux_1_1_x86_64.whl", hash = "sha256:bc604dd77f61c3b6fef06dd73108b390d01699e5535c6d9cf244775cd0bac2f3", size = 4117595, upload-time = "2026-06-25T18:16:34.333Z" }, + { url = "https://files.pythonhosted.org/packages/bc/42/08e91275771c3c72ab643f87e4e4b93d1fcb09769bc0d4df60834443bb68/awscrt-0.35.0-cp313-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:4a7e5f267c494146ca680b1405ad2290b2dcc0821d5f386adcca8aa0bc427c25", size = 3955329, upload-time = "2026-06-25T18:16:39.727Z" }, + { url = "https://files.pythonhosted.org/packages/a1/b9/db8bae837ff816861c06d423c66bfa3e9ee55e975fb26235134c300e04aa/awscrt-0.35.0-cp313-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:9d766a3471f637bf4024fd487a10457e53809e17f3fb2f554bf4b73fad415f58", size = 4252293, upload-time = "2026-06-25T18:16:41.043Z" }, + { url = "https://files.pythonhosted.org/packages/19/9b/37649041dbd8cc06373b9018d882ed42ca6535416e20358c305669cc720b/awscrt-0.35.0-cp313-abi3-musllinux_1_1_aarch64.whl", hash = "sha256:286995476f2f8fd217fbe4c37c7e5bd69953cff03b6ea8d37946537a43208467", size = 3868495, upload-time = "2026-06-25T18:16:42.345Z" }, + { url = "https://files.pythonhosted.org/packages/7d/eb/40c92251e4a17c1fdc2364bab2345791ded4f4da043262c6fadc3f51ed2a/awscrt-0.35.0-cp313-abi3-musllinux_1_1_x86_64.whl", hash = "sha256:177f0f9bbddc3227ede427df0bee9353dbb32fbfdf0d424680400e6d76d11ebc", size = 4111934, upload-time = "2026-06-25T18:16:43.789Z" }, + { url = "https://files.pythonhosted.org/packages/8b/cb/ed8503bcc55c150092278e0019c813cedcb937cad56303e4007e94f43e7a/awscrt-0.35.0-cp313-cp313t-musllinux_1_1_aarch64.whl", hash = "sha256:59976c6063dd79d117d08f93aaddf02448a5b19086b9444c6bdbedf50feff54c", size = 4008528, upload-time = "2026-06-25T18:16:49.675Z" }, + { url = "https://files.pythonhosted.org/packages/8e/5c/461de896c73b2406bb6014422d53f4e0e3a56cc7b856ae0fff570f3878f7/awscrt-0.35.0-cp313-cp313t-musllinux_1_1_x86_64.whl", hash = "sha256:0c933fc94db944f4d2504ffd6bd0c2ef2289a475e0f8896cfea6569be2e4796a", size = 4248623, upload-time = "2026-06-25T18:16:51.227Z" }, +] + [[package]] name = "babel" version = "2.18.0" @@ -3256,6 +3274,8 @@ docs = [ name = "nemo-rl" source = { editable = "." } dependencies = [ + { name = "awscrt", marker = "(platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-fsdp') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-mcore') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-fsdp' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-sglang' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-fsdp') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-mcore') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-vllm') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-fsdp' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-vllm') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-sglang' and extra == 'extra-7-nemo-rl-vllm')" }, + { name = "zstandard", marker = "(platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-fsdp') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-mcore') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-fsdp' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-sglang' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-fsdp') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-mcore') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-vllm') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-fsdp' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-vllm') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-sglang' and extra == 'extra-7-nemo-rl-vllm')" }, { name = "accelerate", marker = "(platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-fsdp') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-mcore') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-fsdp' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-sglang' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-fsdp') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-mcore') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-vllm') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-fsdp' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-vllm') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-sglang' and extra == 'extra-7-nemo-rl-vllm')" }, { name = "blobfile", marker = "(platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-fsdp') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-mcore') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-fsdp' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-sglang' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-fsdp') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-mcore') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-vllm') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-fsdp' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-vllm') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-sglang' and extra == 'extra-7-nemo-rl-vllm')" }, { name = "colored", marker = "(platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-fsdp') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-mcore') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-fsdp' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-sglang') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine != 'aarch64' and platform_machine != 'x86_64' and extra == 'extra-7-nemo-rl-sglang' and extra == 'extra-7-nemo-rl-vllm') or (platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-fsdp') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-mcore') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-automodel' and extra == 'extra-7-nemo-rl-vllm') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-fsdp' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-sglang') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-mcore' and extra == 'extra-7-nemo-rl-vllm') or (sys_platform != 'linux' and extra == 'extra-7-nemo-rl-sglang' and extra == 'extra-7-nemo-rl-vllm')" }, @@ -3424,6 +3444,8 @@ test = [ [package.metadata] requires-dist = [ + { name = "awscrt", specifier = ">=0.35.0" }, + { name = "zstandard" }, { name = "accelerate", specifier = ">=0.26" }, { name = "blobfile" }, { name = "causal-conv1d", marker = "extra == 'automodel'", git = "https://github.com/Dao-AILab/causal-conv1d?rev=4f6ae4e26ae5fe8af9372f8d312ab25cc4595223" }, @@ -7136,6 +7158,20 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/bf/2e/0b49f7e4e53817cfb09a0f6585012b782dfe0b666e8abefcb4fac0570606/z3_solver-4.15.4.0-py3-none-manylinux_2_34_aarch64.whl", hash = "sha256:62c7e9cbdd711932301f29919ad9158de9b2f58b4d281dd259bbcd0a2f408ba1", size = 27226534, upload-time = "2025-10-29T18:11:55.59Z" }, ] + +[[package]] +name = "zstandard" +version = "0.25.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/fd/aa/3e0508d5a5dd96529cdc5a97011299056e14c6505b678fd58938792794b1/zstandard-0.25.0.tar.gz", hash = "sha256:7713e1179d162cf5c7906da876ec2ccb9c3a9dcbdffef0cc7f70c3667a205f0b", size = 711513, upload-time = "2025-09-14T22:15:54.002Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/6d/db/ddb11011826ed7db9d0e485d13df79b58586bfdec56e5c84a928a9a78c1c/zstandard-0.25.0-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:bfc4e20784722098822e3eee42b8e576b379ed72cca4a7cb856ae733e62192ea", size = 5063001, upload-time = "2025-09-14T22:17:31.044Z" }, + { url = "https://files.pythonhosted.org/packages/63/4b/e3678b4e776db00f9f7b2fe58e547e8928ef32727d7a1ff01dea010f3f13/zstandard-0.25.0-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:8e735494da3db08694d26480f1493ad2cf86e99bdd53e8e9771b2752a5c0246a", size = 5547173, upload-time = "2025-09-14T22:17:36.084Z" }, + { url = "https://files.pythonhosted.org/packages/4e/d5/ba05ed95c6b8ec30bd468dfeab20589f2cf709b5c940483e31d991f2ca58/zstandard-0.25.0-cp313-cp313-musllinux_1_1_aarch64.whl", hash = "sha256:3a39c94ad7866160a4a46d772e43311a743c316942037671beb264e395bdd611", size = 5046736, upload-time = "2025-09-14T22:17:37.891Z" }, + { url = "https://files.pythonhosted.org/packages/50/d5/870aa06b3a76c73eced65c044b92286a3c4e00554005ff51962deef28e28/zstandard-0.25.0-cp313-cp313-musllinux_1_1_x86_64.whl", hash = "sha256:172de1f06947577d3a3005416977cce6168f2261284c02080e7ad0185faeced3", size = 5576368, upload-time = "2025-09-14T22:17:40.206Z" }, + { url = "https://files.pythonhosted.org/packages/5d/35/398dc2ffc89d304d59bc12f0fdd931b4ce455bddf7038a0a67733a25f550/zstandard-0.25.0-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:3c83b0188c852a47cd13ef3bf9209fb0a77fa5374958b8c53aaa699398c6bd7b", size = 4954022, upload-time = "2025-09-14T22:17:41.879Z" }, + { url = "https://files.pythonhosted.org/packages/b2/e5/fbd822d5c6f427cf158316d012c5a12f233473c2f9c5fe5ab1ae5d21f3d8/zstandard-0.25.0-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:4f187a0bb61b35119d1926aee039524d1f93aaf38a9916b8c4b78ac8514a0aaf", size = 5360113, upload-time = "2025-09-14T22:17:48.893Z" }, +] [[package]] name = "zipp" version = "3.23.1" From ab626af3580a2d6c55e507d6d317798f36bfe601 Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Fri, 3 Jul 2026 08:36:14 -0700 Subject: [PATCH 02/18] Add zeromq implementation Signed-off-by: Hollow Man --- docs/design-docs/sparse-delta-refit.md | 69 +- examples/configs/grpo_math_1B.yaml | 6 +- nemo_rl/algorithms/grpo.py | 80 ++- nemo_rl/distributed/virtual_cluster.py | 3 +- nemo_rl/models/generation/__init__.py | 3 +- nemo_rl/models/generation/vllm/config.py | 10 +- .../models/generation/vllm/vllm_backend.py | 319 +++++---- .../models/generation/vllm/vllm_generation.py | 50 +- nemo_rl/models/generation/vllm/vllm_worker.py | 386 +++++++---- .../generation/vllm/vllm_worker_async.py | 25 + nemo_rl/models/policy/lm_policy.py | 32 +- .../policy/workers/megatron_policy_worker.py | 100 ++- .../utils/weight_transfer_remote_sparse.py | 630 ++++++++++++++++++ nemo_rl/utils/weight_transfer_s3.py | 122 ---- nemo_rl/utils/weight_transfer_s3_manifest.py | 430 ------------ nemo_rl/utils/weight_transfer_sparse_codec.py | 90 ++- nemo_rl/utils/weight_transfer_zmq.py | 406 +++++++++++ nemo_rl/weight_sync/interfaces.py | 5 +- .../vllm_remote_sparse_weight_synchronizer.py | 221 ++++++ .../vllm_s3_sparse_weight_synchronizer.py | 122 ---- pyrefly.toml | 6 +- tests/unit/algorithms/test_grpo.py | 37 +- .../models/generation/test_vllm_backend.py | 265 +++++--- .../models/generation/test_vllm_generation.py | 260 ++++++-- .../models/policy/test_megatron_worker.py | 62 ++ .../unit/reference_configs/grpo_math_1B.yaml | 6 +- .../test_weight_transfer_remote_sparse.py | 306 +++++++++ .../weight_sync/test_weight_synchronizer.py | 125 ++++ 28 files changed, 2946 insertions(+), 1230 deletions(-) create mode 100644 nemo_rl/utils/weight_transfer_remote_sparse.py delete mode 100644 nemo_rl/utils/weight_transfer_s3.py delete mode 100644 nemo_rl/utils/weight_transfer_s3_manifest.py create mode 100644 nemo_rl/utils/weight_transfer_zmq.py create mode 100644 nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py delete mode 100644 nemo_rl/weight_sync/vllm_s3_sparse_weight_synchronizer.py create mode 100644 tests/unit/utils/test_weight_transfer_remote_sparse.py diff --git a/docs/design-docs/sparse-delta-refit.md b/docs/design-docs/sparse-delta-refit.md index 0cfa9c2c687..8e46292316d 100644 --- a/docs/design-docs/sparse-delta-refit.md +++ b/docs/design-docs/sparse-delta-refit.md @@ -1,8 +1,14 @@ -# S3 Sparse-Delta vLLM Refit +# Remote Sparse-Delta vLLM Refit For non-colocated Megatron policy workers and sync vLLM workers that share the -same checkpoint. Policy workers keep a CPU baseline, upload zstd-compressed -sparse deltas to S3, post receiver manifests, and commit the baseline only after -the global flush succeeds. +same checkpoint. Policy workers keep a CPU baseline and stream zstd-compressed +sparse deltas through either S3 or ZeroMQ. Both transports share export, +encoding, compression, backpressure, receiver apply, and transactional baseline +commit logic. +Payload checksums and transfer-scoped IDs make HTTP retries idempotent. Policy +workers commit their baselines only after every receiver flush succeeds. +On a fresh run, generation starts from the shared checkpoint while policy +workers build the CPU baseline asynchronously; the first transfer follows the +first optimizer step. Resumed runs synchronize before generation. ## Config @@ -11,25 +17,56 @@ backend: vllm colocated: {enabled: false} refit_transport: vllm_s3_sparse delta_compression: - enabled: true dtype: bf16 sparse_bucket_size_bytes: 268435456 vllm_cfg: async_engine: false - expose_http_refit_server: true http_refit_server_port: 8081 http_refit_api_key_env_var: NRL_REFIT_API_KEY ``` -S3 sparse refit requires `kv_cache_dtype: auto`; FP8 KV-cache scale sync is not +Use `refit_transport: vllm_zmq_sparse` and set +`vllm_cfg.zmq_refit_server_port` when Kubernetes needs a stable ZeroMQ target +port. The ZeroMQ service must route TCP traffic from policy workers to the vLLM +relay workers; each relay fans a payload out to every HTTP refit endpoint. On a +flat cluster network, the dynamically reported worker IP can be used directly; +a service mesh is not required. + +Remote sparse refit requires `kv_cache_dtype: auto`; FP8 KV-cache scale sync is not supported. Receiver tensors must have a direct QKV, MoE, Mamba, or generic TP placement plan; transformed and FP8 weights fail before any delta is applied. -The receiver exposes `/nemo-rl/refit/s3-manifest`, with -`http_refit_api_key_env_var` auth when configured. Export chunks are capped by -`NRL_REFIT_S3_EXPORT_CHUNK_BYTES` and -`delta_compression.sparse_bucket_size_bytes`. Set `NRL_REFIT_S3_BUCKET` and, -when needed, `NRL_REFIT_S3_REGION` or `NRL_REFIT_S3_PREFIX`. AWS CRT performs -multipart transfer automatically; encode and end-to-end pipeline concurrency -default from available CPU cores and can be fixed with -`NRL_REFIT_S3_ENCODE_WORKERS` and `NRL_REFIT_S3_UPLOAD_WORKERS`. Track `REFIT_S3_TIMING`, -`REFIT_RECEIVER_TIMING`, and `REFIT_S3_GLOBAL_FLUSH` in cluster runs. +The HTTP endpoints and ZeroMQ producer relay use +`http_refit_api_key_env_var` auth when configured. + +For S3, set `NRL_REFIT_S3_BUCKET` and, when needed, +`NRL_REFIT_S3_REGION` or `NRL_REFIT_S3_PREFIX`. AWS CRT performs multipart +transfer automatically. Tune export, encode, and transfer concurrency with +`NRL_REFIT_S3_EXPORT_CHUNK_BYTES`, `NRL_REFIT_S3_ENCODE_WORKERS`, and +`NRL_REFIT_S3_UPLOAD_WORKERS`. + +For ZeroMQ, tune the same stages with `NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES`, +`NRL_REFIT_ZMQ_ENCODE_WORKERS`, and `NRL_REFIT_ZMQ_SEND_WORKERS`. Relay +concurrency is controlled by `NRL_REFIT_ZMQ_RELAY_PAYLOAD_WORKERS` and +`NRL_REFIT_ZMQ_RELAY_FANOUT_WORKERS`; fanout defaults to 32 workers to preserve +HTTP keep-alive reuse and avoid receiver contention. Track `REFIT_S3_TIMING` or +`REFIT_ZMQ_TIMING`, `REFIT_RECEIVER_TIMING`, and `REFIT_*_GLOBAL_COMMIT` in +cluster runs. + +Set `NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD` to a small positive value to verify +deterministic transmitted-delta samples after placement. Each receiver snapshots +only those target elements before apply and compares `post - pre` with the +placement- and dtype-adjusted transmitted delta, so an existing absolute weight +offset cannot contaminate the metric. `REFIT_*_DELTA_VERIFY` reports candidate +and applied sample counts, exact mismatches, tolerance-gated mismatches +(`rtol=1e-6`, `atol=1e-8`), and mean/max absolute delta error. A gated mismatch +aborts the transaction before baseline commit. Successful commits advance the +CPU baseline to the quantized value applied by the receiver, so later deltas +compensate compression residuals instead of accumulating drift. + +`REFIT_*_DELTA_CHANGE` reports changed and total exported element counts plus +their model-wide percentage. The codec accumulates these counters while it is +already finding sparse locations; it does not perform another tensor scan. +When a training logger is enabled, the same values are emitted to W&B and +TensorBoard under `refit/delta/*`, `refit/delta_verify/*`, and +`refit/transfer/*`. End-to-end refit latency remains available as +`timing/train/prepare_for_generation/transfer_and_update_weights`. diff --git a/examples/configs/grpo_math_1B.yaml b/examples/configs/grpo_math_1B.yaml index 1a742b53cf3..7eb30cf5c89 100644 --- a/examples/configs/grpo_math_1B.yaml +++ b/examples/configs/grpo_math_1B.yaml @@ -339,8 +339,8 @@ policy: top_k: null stop_token_ids: null stop_strings: null - refit_transport: null # Set to "vllm_s3_sparse" to use S3 sparse-delta refit. - delta_compression: null # S3 sparse-delta refit config; null uses the existing refit path. + refit_transport: null # Set to "vllm_s3_sparse" or "vllm_zmq_sparse" for remote sparse-delta refit. + delta_compression: null # Remote sparse-delta config; null uses the existing refit path. mcore_generation_config: async_engine: false max_model_len: ${policy.max_total_sequence_length} # Engine-side max sequence length. @@ -378,9 +378,9 @@ policy: num_first_layers_in_bf16: 0 enable_vllm_metrics_logger: true # Set to true to enable vLLM internal metrics logger, turn off for better performance vllm_metrics_logger_interval: 0.5 # Interval in seconds to collect vLLM logger metrics - expose_http_refit_server: false # Start the internal sparse-delta refit endpoint on vLLM workers. http_refit_api_key_env_var: null # Optional env var containing the internal refit API key. http_refit_server_port: null # Optional fixed port for Kubernetes targetPorts. + zmq_refit_server_port: null # Optional fixed ZeroMQ relay port for Kubernetes targetPorts. vllm_kwargs: {} colocated: # true: generation shares training GPUs diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index fe13717eb53..afa5938d390 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -118,8 +118,8 @@ from nemo_rl.utils.nsys import maybe_gpu_profile_step from nemo_rl.utils.timer import TimeoutChecker, Timer from nemo_rl.utils.venvs import create_local_venv_on_each_node -from nemo_rl.weight_sync.vllm_s3_sparse_weight_synchronizer import ( - VllmS3SparseWeightSynchronizer, +from nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer import ( + VllmRemoteSparseWeightSynchronizer, ) # =============================================================================== @@ -910,9 +910,8 @@ def _spinup_nemo_gym(base_urls, model_name): # vllm model loading prefers clean environment, initialize policy_generation before policy in colocated mode backend = generation_config["backend"] generation_config["model_name"] = policy_config["model_name"] # Needed for vLLM - use_vllm_s3_sparse_refit = ( - backend == "vllm" - and generation_config.get("refit_transport") == "vllm_s3_sparse" + refit_transport = ( + generation_config.get("refit_transport") if backend == "vllm" else None ) # Dictionary to store worker initialization timing stats for logging @@ -1103,27 +1102,21 @@ def initialize_generation_with_policy( # vLLM generation: setup config, then initialize with policy generation_config = cast(VllmConfig, generation_config) vllm_cfg = generation_config["vllm_cfg"] - refit_transport = generation_config.get("refit_transport", None) - if refit_transport not in (None, "vllm_s3_sparse"): + if refit_transport not in (None, "vllm_s3_sparse", "vllm_zmq_sparse"): raise ValueError(f"Unsupported vLLM refit transport {refit_transport!r}.") - if use_vllm_s3_sparse_refit: - delta_config = generation_config.get("delta_compression") + if refit_transport is not None: if ( colocated_inference or not policy_config["megatron_cfg"]["enabled"] - or vllm_cfg["async_engine"] or vllm_cfg["precision"] == "fp8" or vllm_cfg["kv_cache_dtype"].startswith("fp8") - or not delta_config - or not delta_config.get("enabled") + or not generation_config.get("delta_compression") or generation_config.get("quant_cfg") or generation_config.get("real_quant") - or not vllm_cfg.get("expose_http_refit_server") ): raise ValueError( - "vllm_s3_sparse requires a non-colocated Megatron policy, " - "synchronous BF16/FP16 vLLM, delta compression, an unquantized " - "rollout, and the refit HTTP server." + f"{refit_transport} requires a non-colocated Megatron policy, " + "BF16/FP16 vLLM, delta compression, and an unquantized rollout." ) if generation_config["vllm_cfg"]["precision"] == "fp8": @@ -1263,7 +1256,7 @@ def init_vllm_then_policy(): policy.print_node_ip_and_gpu_id() # if it is not colocated inference, initialize collective communication for update weights - if not colocated_inference and not use_vllm_s3_sparse_refit: + if not colocated_inference and refit_transport is None: t0 = time.perf_counter() ip, port = train_cluster.get_master_address_and_port() print(f"Using ip: {ip}, port: {port} for collective communication", flush=True) @@ -1302,18 +1295,19 @@ def init_vllm_then_policy(): ray.get(futures_train + futures_inference) worker_init_timing_metrics["collective_init_time_s"] = time.perf_counter() - t0 - if use_vllm_s3_sparse_refit: + if refit_transport is not None: t0 = time.perf_counter() assert isinstance(policy_generation, VllmGeneration) - policy_generation.weight_synchronizer = VllmS3SparseWeightSynchronizer( + policy_generation.weight_synchronizer = VllmRemoteSparseWeightSynchronizer( policy, policy_generation, + transport=refit_transport.removeprefix("vllm_").removesuffix("_sparse"), api_key_env_var=generation_config["vllm_cfg"].get( "http_refit_api_key_env_var" ), ) policy_generation.weight_synchronizer.init_communicator() - worker_init_timing_metrics["s3_sparse_refit_init_time_s"] = ( + worker_init_timing_metrics[f"{refit_transport}_init_time_s"] = ( time.perf_counter() - t0 ) else: @@ -2079,7 +2073,7 @@ def refit_policy_generation( _refit_buffer_size_gb: Optional[float] = None, timer: Optional[Timer] = None, kv_scales: Optional[dict[str, float]] = None, -) -> None: +) -> dict[str, float]: """Refit the policy generation interface with the latest policy weights. Args: @@ -2089,16 +2083,21 @@ def refit_policy_generation( the buffer size is computed from remaining memory. timer: Optional Timer used to time the prepare/transfer/update phase kv_scales: Optional dictionary of KV cache scales for FP8 quantization. + + Returns: + Scalar metrics reported by the selected weight synchronizer. """ if ( isinstance(policy_generation, VllmGeneration) and policy_generation.weight_synchronizer is not None ): - policy_generation.weight_synchronizer.sync_weights( - timer=timer, - kv_scales=kv_scales, + return ( + policy_generation.weight_synchronizer.sync_weights( + timer=timer, + kv_scales=kv_scales, + ) + or {} ) - return # Megatron generation backend needs explicit suspend/resume around refits. if isinstance(policy_generation, MegatronGeneration): @@ -2200,6 +2199,16 @@ def refit_policy_generation( if isinstance(policy_generation, MegatronGeneration): policy_generation.resume_after_refit() + return {} + + +def _initial_policy_generation_stale( + policy_generation: GenerationInterface, completed_steps: int +) -> bool: + """Skip a fresh run's redundant sync when the synchronizer is already current.""" + synchronizer = getattr(policy_generation, "weight_synchronizer", None) + return completed_steps > 0 or synchronizer is None or synchronizer.is_stale + def _log_mixed_rewards_and_advantages_information( logger: Logger, @@ -2389,7 +2398,6 @@ def grpo_train( isinstance(policy_generation, MegatronGeneration) and master_config.policy["generation"]["colocated"]["enabled"] ) - POLICY_GENERATION_STALE = True # tracks if generation needs a refit before running assert policy_generation is not None # Check if we need to sync KV cache scales @@ -2399,6 +2407,9 @@ def grpo_train( # common config/state times current_step = grpo_save_state["current_step"] # current step within an epoch total_steps = grpo_save_state["total_steps"] # total steps across all epochs + POLICY_GENERATION_STALE = _initial_policy_generation_stale( + policy_generation, total_steps + ) max_num_steps = master_config.grpo[ "max_num_steps" ] # max number of steps to train for @@ -2466,6 +2477,7 @@ def grpo_train( batch_cache: BatchedDataDict[DatumSpec] = None # This is the number of batches we processed so far at each step to generate responses whose std is non-zero. Maximum threshold is set by dynamic_sampling_max_gen_batches. Used in the case of dynamic sampling. dynamic_sampling_num_gen_batches = 0 + refit_metrics: dict[str, float] = {} # Run grpo/dapo training loop (single-turn) for batch in wrapped_dataloader: @@ -2545,7 +2557,7 @@ def grpo_train( calibration_data, include_q=True )["layers"] - refit_policy_generation( + refit_metrics = refit_policy_generation( policy, policy_generation, colocated_inference, @@ -3000,7 +3012,7 @@ def grpo_train( ): memory_tracker.snapshot_start_of_stage("Validation", dir()) if NEED_REFIT and POLICY_GENERATION_STALE: - refit_policy_generation( + refit_metrics = refit_policy_generation( policy, policy_generation, colocated_inference, @@ -3352,6 +3364,8 @@ def grpo_train( train_results, metrics, timing_metrics, master_config ) + if refit_metrics: + logger.log_metrics(refit_metrics, total_steps + 1, prefix="refit") logger.log_metrics(metrics, total_steps + 1, prefix="train") logger.log_metrics( performance_metrics, total_steps + 1, prefix="performance" @@ -3367,6 +3381,7 @@ def grpo_train( # Reset the batch and set dynamic_sampling_num_gen_batches to 0 batch_cache = None dynamic_sampling_num_gen_batches = 0 + refit_metrics = {} # Clear mem memory_tracker.snapshot_start_of_stage("After CPU memory clear", dir()) @@ -3683,11 +3698,11 @@ def async_grpo_train( isinstance(policy_generation, MegatronGeneration) and master_config.policy["generation"]["colocated"]["enabled"] ) - POLICY_GENERATION_STALE = True assert policy_generation is not None # Training state step = grpo_save_state["current_step"] + POLICY_GENERATION_STALE = _initial_policy_generation_stale(policy_generation, step) weight_version = step # Tracks refitted weight versions consumed_samples = grpo_save_state["consumed_samples"] total_valid_tokens = grpo_save_state.get( @@ -3951,6 +3966,7 @@ def async_grpo_train( # Main training loop try: while step < master_config.grpo["max_num_steps"]: + refit_metrics: dict[str, float] = {} print( f"\n{'=' * 25} Step {step + 1}/{master_config.grpo['max_num_steps']} {'=' * 25}" ) @@ -4343,7 +4359,7 @@ def async_grpo_train( # Only the actual refit/weight transfer should be counted as weight_sync print("🔄 Performing policy generation refit...") with timer.time("weight_sync"): - refit_policy_generation( + refit_metrics = refit_policy_generation( policy, policy_generation, colocated_inference, @@ -4374,7 +4390,7 @@ def async_grpo_train( trajectory_collector.pause.remote() if NEED_REFIT and POLICY_GENERATION_STALE: - refit_policy_generation( + refit_metrics = refit_policy_generation( policy, policy_generation, colocated_inference ) POLICY_GENERATION_STALE = False @@ -4701,6 +4717,8 @@ def async_grpo_train( merged_efficiency, total_wall_time, step + 1 ) + if refit_metrics: + logger.log_metrics(refit_metrics, step + 1, prefix="refit") logger.log_metrics(performance_metrics, step + 1, prefix="performance") logger.log_metrics(metrics, step + 1, prefix="train") logger.log_metrics(efficiency_loggable, step + 1, prefix="") diff --git a/nemo_rl/distributed/virtual_cluster.py b/nemo_rl/distributed/virtual_cluster.py index c4e80445070..52926991745 100644 --- a/nemo_rl/distributed/virtual_cluster.py +++ b/nemo_rl/distributed/virtual_cluster.py @@ -1095,4 +1095,5 @@ def __del__(self) -> None: the cluster is lost due to leaving a function scope. It's always recommended that the user calls shutdown(). """ - self.shutdown() + if not sys.is_finalizing(): + self.shutdown() diff --git a/nemo_rl/models/generation/__init__.py b/nemo_rl/models/generation/__init__.py index 26c82e50a2e..f208a40a945 100644 --- a/nemo_rl/models/generation/__init__.py +++ b/nemo_rl/models/generation/__init__.py @@ -47,7 +47,8 @@ def configure_generation_config( # set load_format config["vllm_cfg"]["load_format"] = ( "auto" - if is_eval or config.get("refit_transport") == "vllm_s3_sparse" + if is_eval + or config.get("refit_transport") in ("vllm_s3_sparse", "vllm_zmq_sparse") else "dummy" ) speculative_config = config.get("vllm_kwargs", {}).get("speculative_config") diff --git a/nemo_rl/models/generation/vllm/config.py b/nemo_rl/models/generation/vllm/config.py index 1e14f572ae2..0f4a86b687e 100644 --- a/nemo_rl/models/generation/vllm/config.py +++ b/nemo_rl/models/generation/vllm/config.py @@ -41,12 +41,12 @@ class VllmSpecificArgs(TypedDict): # Exposing vLLM as a server is useful in instances where the multi-turn rollout is performed with utilities outside of NeMo RL, but the user still wants to take advantage of the refit logic in NeMo RL that keeps the policy and generation up to date. # Currently it will expose the /tokenize and /v1/chat/completions endpoints. Later on we may expose /v1/completions or /v1/responses. expose_http_server: NotRequired[bool] - # Internal trusted endpoint for sparse delta refit payloads. - expose_http_refit_server: NotRequired[bool] # Environment variable containing the internal refit API key. http_refit_api_key_env_var: NotRequired[str | None] # Fixed internal refit endpoint port for stable Kubernetes targetPorts. http_refit_server_port: NotRequired[int | None] + # Fixed ZeroMQ relay port for stable Kubernetes targetPorts. + zmq_refit_server_port: NotRequired[int | None] # These kwargs are passed to the vllm.LLM HTTP server Chat Completions endpoint config. Typically this will include things like tool parser, chat template, etc http_server_serving_chat_kwargs: NotRequired[dict[str, Any]] # Miscellaneous top level vLLM HTTP server arguments. @@ -61,17 +61,15 @@ class VllmSpecificArgs(TypedDict): class VllmDeltaCompressionConfig(TypedDict): - enabled: bool dtype: DeltaCompressionDType sparse_bucket_size_bytes: int - baseline_mmap_dir: NotRequired[str | None] class VllmConfig(GenerationConfig): vllm_cfg: VllmSpecificArgs vllm_kwargs: NotRequired[dict[str, Any]] - # Null uses the existing NCCL refit; "vllm_s3_sparse" uses S3 sparse deltas. - refit_transport: NotRequired[Literal["vllm_s3_sparse"] | None] + # Null uses NCCL; remote sparse refit supports S3 or ZeroMQ value planes. + refit_transport: NotRequired[Literal["vllm_s3_sparse", "vllm_zmq_sparse"] | None] delta_compression: NotRequired[VllmDeltaCompressionConfig | None] # quantization config diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index f975c982d8b..0437767a309 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -118,10 +118,13 @@ class _SparseDeltaTargetPlan: class VllmInternalWorkerExtension: state_dict_info: dict[str, Any] | None = None _direct_sparse_delta_targets: dict[str, torch.Tensor] | None = None - _direct_sparse_delta_modules: dict[str, torch.nn.Module] | None = None _direct_sparse_delta_plan_cache: dict[str, _SparseDeltaTargetPlan | None] | None = ( None ) + _direct_sparse_delta_verification: ( + list[tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]] | None + ) = None + _direct_sparse_delta_verification_candidates = 0 def bind_numa(self) -> bool: """Pin this TP worker to its GPU's NUMA-local CPUs/memory. @@ -203,8 +206,9 @@ def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: """ 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 self._direct_sparse_delta_targets = None - self._direct_sparse_delta_modules = None self._direct_sparse_delta_plan_cache = None + self._direct_sparse_delta_verification = [] + self._direct_sparse_delta_verification_candidates = 0 def _process_weights_after_loading( self, @@ -404,16 +408,32 @@ def _apply_sparse_weight_deltas( metadata: list[dict[str, Any]], ) -> None: """Apply sparse deltas directly after validating every target plan.""" - if self._direct_sparse_delta_uses_loader_transform(): + architectures = self.model_runner.vllm_config.model_config.architectures + from nemo_rl.models.generation.vllm.quantization import fp8 + + if {"GptOssForCausalLM", "Gemma3ForConditionalGeneration"} & set( + architectures + ) or fp8.is_fp8_model(self.model_runner.vllm_config): raise RuntimeError( "Direct sparse delta refit does not support transformed or FP8 weights." ) - targets = self._direct_sparse_delta_target_map() + if self._direct_sparse_delta_targets is None: + model = self.model_runner.model + self._direct_sparse_delta_targets = dict(model.named_parameters()) | dict( + model.named_buffers() + ) + targets = self._direct_sparse_delta_targets raw_locations, raw_values = payload_tensors + plan_cache = self._direct_sparse_delta_plan_cache + if plan_cache is None: + plan_cache = self._direct_sparse_delta_plan_cache = {} plans = [] for item in metadata: - plan = self._cached_direct_sparse_delta_target_plan(item, targets) + name = str(item["name"]) + if name not in plan_cache: + plan_cache[name] = self._direct_sparse_delta_target_plan(item, targets) + plan = plan_cache[name] if plan is None: raise RuntimeError( f"No direct sparse delta target plan for {item['name']!r}." @@ -423,9 +443,42 @@ def _apply_sparse_weight_deltas( with torch.no_grad(): for item, plan in plans: target = plan.target + verification_locations = item.get("verification_locations", []) + self._direct_sparse_delta_verification_candidates += len( + verification_locations + ) if target is None: continue + if verification_locations and not plan.log_delta_transform: + sample_locations, sample_deltas = ( + self._local_sparse_delta_update_inputs( + torch.tensor(verification_locations, device=target.device), + torch.tensor( + item["verification_deltas"], + device=target.device, + dtype=target.dtype, + ), + plan, + ) + ) + if sample_locations.numel(): + before = target.data.view(-1).index_select(0, sample_locations) + expected_delta = ( + before + sample_deltas + ).float() - before.float() + verification = self._direct_sparse_delta_verification + if verification is None: + verification = self._direct_sparse_delta_verification = [] + verification.append( + ( + target, + sample_locations, + before.float(), + expected_delta, + ) + ) + value_start = int(item["value_start"]) value_end = int(item["value_end"]) values = raw_values[value_start:value_end].to( @@ -439,69 +492,37 @@ def _apply_sparse_weight_deltas( target.data.view(-1).narrow(0, range_start, range_count).add_( values ) - continue - - locations = sparse_codec.sparse_locations_for_item( - item, - raw_locations, - device=target.device, - ) - locations, values = self._local_sparse_delta_update_inputs( - locations, - values, - plan, - ) - if locations.numel(): - target_flat = target.data.view(-1) - if plan.log_delta_transform: - current = target_flat.index_select(0, locations) - updated = current * values.float().exp().to(dtype=current.dtype) - target_flat.index_copy_(0, locations, updated) - else: - target_flat.index_add_(0, locations, values) - - def _direct_sparse_delta_target_map(self) -> dict[str, torch.Tensor]: - if self._direct_sparse_delta_targets is None: - self._direct_sparse_delta_targets = dict( - self.model_runner.model.named_parameters() - ) - self._direct_sparse_delta_targets.update( - self.model_runner.model.named_buffers() - ) - return self._direct_sparse_delta_targets - - def _direct_sparse_delta_named_module( + else: + locations = sparse_codec.sparse_locations_for_item( + item, + raw_locations, + device=target.device, + ) + locations, values = self._local_sparse_delta_update_inputs( + locations, + values, + plan, + ) + if locations.numel(): + target_flat = target.data.view(-1) + if plan.log_delta_transform: + current = target_flat.index_select(0, locations) + updated = current * values.float().exp().to( + dtype=current.dtype + ) + target_flat.index_copy_(0, locations, updated) + else: + target_flat.index_add_(0, locations, values) + + def _direct_sparse_delta_module( self, + target: torch.Tensor, module_name: str, - ) -> torch.nn.Module | None: - if self._direct_sparse_delta_modules is None: - self._direct_sparse_delta_modules = dict( - self.model_runner.model.named_modules() - ) - return self._direct_sparse_delta_modules.get(module_name) - - def _direct_sparse_delta_uses_loader_transform(self) -> bool: - architectures = self.model_runner.vllm_config.model_config.architectures - if {"GptOssForCausalLM", "Gemma3ForConditionalGeneration"} & set(architectures): - return True - - from nemo_rl.models.generation.vllm.quantization import fp8 - - return fp8.is_fp8_model(self.model_runner.vllm_config) - - def _cached_direct_sparse_delta_target_plan( - self, - item: dict[str, Any], - targets: dict[str, torch.Tensor], - ) -> _SparseDeltaTargetPlan | None: - name = str(item["name"]) - if self._direct_sparse_delta_plan_cache is None: - self._direct_sparse_delta_plan_cache = {} - if name not in self._direct_sparse_delta_plan_cache: - self._direct_sparse_delta_plan_cache[name] = ( - self._direct_sparse_delta_target_plan(item, targets) - ) - return self._direct_sparse_delta_plan_cache[name] + ) -> Any: + loader = getattr(target, "weight_loader", None) + return getattr( + loader, "__self__", None + ) or self.model_runner.model.get_submodule(module_name) def _direct_sparse_delta_target_plan( self, @@ -511,8 +532,9 @@ def _direct_sparse_delta_target_plan( name = str(item["name"]) if name.startswith("mtp."): return _SparseDeltaTargetPlan(target=None) - target_name = self._map_direct_sparse_delta_name(name) - if target_name is None: + mapper = getattr(self.model_runner.model, "hf_to_vllm_mapper", None) + target_name = cast(Any, mapper)._map_name(name) if mapper is not None else name + if target_name is None or target_name.startswith("draft."): return None if ".mixer." in target_name: mamba_plan = self._direct_sparse_delta_mamba2_plan( @@ -526,6 +548,12 @@ def _direct_sparse_delta_target_plan( return self._direct_sparse_delta_qkv_plan(item, target_name, targets) if _EXPERT_WEIGHT_RE.match(target_name): return self._direct_sparse_delta_expert_plan(item, target_name, targets) + if any(f".{candidate}_proj." in target_name for candidate in ("gate", "up")): + merged_plan = self._direct_sparse_delta_merged_column_plan( + item, target_name, targets + ) + if merged_plan is not None: + return merged_plan target = targets.get(target_name) if target is None: @@ -548,11 +576,7 @@ def _direct_sparse_delta_qkv_plan( if target is None: return None output_dim = int(cast(Any, target).output_dim) % target.ndim - module = cast( - Any, - getattr(getattr(target, "weight_loader", None), "__self__", None) - or self._direct_sparse_delta_named_module(packed_name.rsplit(".", 1)[0]), - ) + module = self._direct_sparse_delta_module(target, packed_name.rsplit(".", 1)[0]) shard_offset = int(module._get_shard_offset_mapping(shard_id)) shard_size = int(module._get_shard_size_mapping(shard_id)) shard_rank = int(module.tp_rank) @@ -572,6 +596,52 @@ def _direct_sparse_delta_qkv_plan( target_offset=shard_offset * target.stride(output_dim), ) + def _direct_sparse_delta_merged_column_plan( + self, + item: dict[str, Any], + target_name: str, + targets: dict[str, torch.Tensor], + ) -> _SparseDeltaTargetPlan | None: + projection = next( + candidate + for candidate in ("gate", "up") + if f".{candidate}_proj." in target_name + ) + shard_id = 0 if projection == "gate" else 1 + packed_name = target_name.replace(f".{projection}_proj.", ".gate_up_proj.", 1) + target = targets.get(packed_name) + output_dim = getattr(target, "output_dim", None) + if target is None or not isinstance(output_dim, int): + return None + + output_dim %= target.ndim + module = self._direct_sparse_delta_module(target, packed_name.rsplit(".", 1)[0]) + output_sizes = tuple(int(size) for size in module.output_sizes) + tp_size = int(module.tp_size) + source_shape = tuple(item["shape"]) + if ( + shard_id >= len(output_sizes) + or tp_size < 1 + or output_sizes[shard_id] % tp_size + or output_dim >= len(source_shape) + or source_shape[output_dim] != output_sizes[shard_id] + ): + return None + + shard_size = output_sizes[shard_id] // tp_size + target_start = sum(output_sizes[:shard_id]) // tp_size + if target.shape[output_dim] < target_start + shard_size: + return None + shard_start = int(module.tp_rank) * shard_size + return self._make_sparse_delta_target_plan( + target, + source_shape=source_shape, + shard_dim=output_dim, + shard_start=shard_start, + shard_size=shard_size, + target_offset=target_start * target.stride(output_dim), + ) + def _direct_sparse_delta_mamba2_plan( self, item: dict[str, Any], @@ -600,7 +670,7 @@ def _direct_sparse_delta_mamba2_plan( return None mixer_name = target_name.split(".mixer.", 1)[0] + ".mixer" - attrs = cast(Any, self._direct_sparse_delta_named_module(mixer_name)) + attrs = cast(Any, self.model_runner.model.get_submodule(mixer_name)) tp_size = int(attrs.tp_size) if tp_size <= 1: return None @@ -616,7 +686,7 @@ def _direct_sparse_delta_mamba2_plan( if remainder or group_size <= 0 or extra_group_size < 0: return None tp_rank = int( - getattr(target, "tp_rank", self._direct_sparse_delta_tp_rank(tp_size)) + getattr(target, "tp_rank", int(getattr(self, "rank", 0)) % tp_size) ) intermediate = (intermediate_size, 0, False) group = (groups_ssm_state_size, extra_group_size, extra_group_size > 0) @@ -677,10 +747,8 @@ def _direct_sparse_delta_expert_plan( target = targets.get(packed_name) if target is None: return None - module_attrs = cast( - Any, - getattr(getattr(target, "weight_loader", None), "__self__", None) - or self._direct_sparse_delta_named_module(packed_name.rsplit(".", 1)[0]), + module_attrs = self._direct_sparse_delta_module( + target, packed_name.rsplit(".", 1)[0] ) if shard_id == "w3" and not module_attrs.moe_config.is_act_and_mul: shard_id = "w1" @@ -815,7 +883,7 @@ def _direct_sparse_delta_shard_plan( if source_shape[shard_dim] > shard_size * tp_size: continue tp_rank = int( - getattr(target, "tp_rank", self._direct_sparse_delta_tp_rank(tp_size)) + getattr(target, "tp_rank", int(getattr(self, "rank", 0)) % tp_size) ) plan = self._make_sparse_delta_target_plan( target=target, @@ -905,22 +973,6 @@ def _local_sparse_delta_update_inputs( local_locations.add_(coord * target_stride) return local_locations, selected_values - def _direct_sparse_delta_tp_rank(self, tp_size: int) -> int: - if tp_size <= 1: - return 0 - rank = int(getattr(self, "rank", 0)) - return rank % tp_size - - def _map_direct_sparse_delta_name(self, name: str) -> str | None: - mapper = getattr(self.model_runner.model, "hf_to_vllm_mapper", None) - if mapper is not None: - name = cast(Any, mapper)._map_name(name) - if name is None: - return None - if name.startswith("draft."): - return None - return name - @wrap_with_nvtx_name("vllm_internal_worker_extension/update_weights_via_ipc_zmq") def update_weights_via_ipc_zmq(self) -> bool: """Receive and update model weights via ZMQ IPC socket. @@ -1054,17 +1106,13 @@ def update_weights_from_collective(self) -> bool: def update_weights_from_serialized_sparse_payload( self, serialized_payload: bytes, - synchronize: bool = True, ) -> dict[str, Any]: - """Apply one serialized sparse-delta payload received from S3.""" - return self._load_and_apply_sparse_payload( - io.BytesIO(serialized_payload), synchronize - ) + """Apply one serialized sparse-delta payload.""" + return self._load_and_apply_sparse_payload(io.BytesIO(serialized_payload)) def _load_and_apply_sparse_payload( self, source: str | io.BytesIO, - synchronize: bool, ) -> dict[str, Any]: started = time.perf_counter() payload = cast( @@ -1076,7 +1124,7 @@ def _load_and_apply_sparse_payload( ), ) deserialize_s = time.perf_counter() - started - result = self._apply_sparse_request(payload, synchronize=synchronize) + result = self._apply_sparse_request(payload) result["receiver_deserialize_s"] = deserialize_s result["receiver_total_s"] = time.perf_counter() - started return result @@ -1087,20 +1135,15 @@ def _load_and_apply_sparse_payload( def update_weights_from_sparse_payload_files( self, *payload_paths: str, - synchronize: bool = True, ) -> dict[str, Any]: - """Apply sparse payloads in FIFO order with one final sync.""" - if not payload_paths: - raise ValueError("A sparse refit batch must contain at least one payload.") + """Apply sparse payloads in FIFO order.""" started = time.perf_counter() deserialize_s = 0.0 sparse_apply_s = 0.0 for path in payload_paths: - result = self._load_and_apply_sparse_payload(path, synchronize=False) + result = self._load_and_apply_sparse_payload(path) deserialize_s += float(result["receiver_deserialize_s"]) sparse_apply_s += float(result["receiver_sparse_apply_s"]) - if synchronize and torch.cuda.is_available(): - torch.cuda.synchronize(self.device) return { "ok": True, "receiver_deserialize_s": deserialize_s, @@ -1111,8 +1154,6 @@ def update_weights_from_sparse_payload_files( def _apply_sparse_request( self, payload: sparse_codec.TensorPayload, - *, - synchronize: bool, ) -> dict[str, Any]: locations, values, metadata = payload @@ -1120,18 +1161,64 @@ def _apply_sparse_request( self._apply_sparse_weight_deltas((locations, values), metadata) sparse_apply_s = time.perf_counter() - sparse_started - if synchronize and torch.cuda.is_available(): - torch.cuda.synchronize(self.device) return { "ok": True, "receiver_sparse_apply_s": sparse_apply_s, } - def synchronize_device(self) -> dict[str, Any]: + def synchronize_device(self) -> None: """Synchronize this vLLM worker's CUDA device after deferred refit applies.""" if torch.cuda.is_available(): torch.cuda.synchronize(self.device) - return {"ok": True} + + def finish_sparse_delta_refit(self) -> dict[str, Any]: + """Synchronize and compare bounded producer samples with applied weights.""" + self.synchronize_device() + verification = self._direct_sparse_delta_verification or [] + candidates = self._direct_sparse_delta_verification_candidates + self._direct_sparse_delta_verification = [] + self._direct_sparse_delta_verification_candidates = 0 + if not verification: + return { + "ok": True, + "verification_candidates": candidates, + "verification_samples": 0, + "verification_exact_mismatches": 0, + "verification_mismatches": 0, + "verification_abs_sum": 0.0, + "verification_max_abs": 0.0, + } + + with torch.no_grad(): + actual_delta = torch.cat( + [ + target.data.view(-1).index_select(0, locations).float() - before + for target, locations, before, _ in verification + ] + ) + expected_delta = torch.cat([expected for _, _, _, expected in verification]) + difference = (actual_delta - expected_delta).abs() + exact_mismatches = actual_delta.ne(expected_delta) + mismatches = ~torch.isclose( + actual_delta, expected_delta, rtol=1e-6, atol=1e-8 + ) + stats = torch.stack( + ( + difference.sum(), + difference.max(), + exact_mismatches.sum().float(), + mismatches.sum().float(), + ) + ).cpu() + return { + "ok": True, + "verification_candidates": candidates, + "verification_samples": actual_delta.numel(), + "verification_exact_mismatches": int(stats[2]), + "verification_mismatches": int(stats[3]), + "verification_abs_sum": float(stats[0]), + "verification_max_abs": float(stats[1]), + } def cleanup(self) -> None: """Shutdown and cleanup resources.""" diff --git a/nemo_rl/models/generation/vllm/vllm_generation.py b/nemo_rl/models/generation/vllm/vllm_generation.py index 058347c9f66..e51942d1259 100644 --- a/nemo_rl/models/generation/vllm/vllm_generation.py +++ b/nemo_rl/models/generation/vllm/vllm_generation.py @@ -15,6 +15,7 @@ import asyncio import logging import os +import sys import warnings from collections import defaultdict from typing import ( @@ -905,6 +906,8 @@ def finish_generation(self, *args: Any, **kwargs: Any) -> bool: def shutdown(self) -> bool: """Shut down all vLLM workers and clean up resources.""" try: + if self.weight_synchronizer is not None: + self.weight_synchronizer.shutdown() # Use the worker group's shutdown method with the worker's cleanup method return self.worker_group.shutdown(cleanup_method="shutdown") except Exception as e: @@ -913,33 +916,45 @@ def shutdown(self) -> bool: def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: """Prepare the info for refit.""" - # Choose the appropriate method based on async_engine setting method_name = ( "prepare_refit_info_async" if self.cfg["vllm_cfg"]["async_engine"] else "prepare_refit_info" ) - - # Use run_all_workers_single_data to send data to all workers - futures = self.worker_group.run_all_workers_single_data( - method_name, - state_dict_info=state_dict_info, - run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], - ) - - # Wait for all futures to complete - ray.get(futures) + self._run_refit_workers(method_name, state_dict_info=state_dict_info) def report_refit_server_base_urls(self) -> list[str]: """Return base URLs for vLLM workers exposing sparse refit endpoints.""" + return [ + url + for url in self._run_refit_workers("report_refit_server_base_url") + if url + ] + + def _run_refit_workers(self, method_name: str, **kwargs: Any) -> list[Any]: if not self.worker_group or not self.worker_group.workers: raise RuntimeError("Worker group is not initialized") - - futures = self.worker_group.run_all_workers_single_data( - "report_refit_server_base_url", - run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], + return ray.get( + self.worker_group.run_all_workers_single_data( + method_name, + run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], + **kwargs, + ) ) - return [url for url in ray.get(futures) if url] + + def start_zmq_sparse_refit_relays(self, refit_urls: list[str]) -> list[str]: + """Start one ZeroMQ relay per vLLM replica and return TCP addresses.""" + return [ + address + for address in self._run_refit_workers( + "start_zmq_sparse_refit_relay", refit_urls=refit_urls + ) + if address + ] + + def stop_zmq_sparse_refit_relays(self) -> None: + if self.worker_group and self.worker_group.workers: + self._run_refit_workers("stop_zmq_sparse_refit_relay") def update_weights_via_ipc_zmq(self) -> list[ray.ObjectRef]: """Update weights of the policy using IPC handles via ZMQ socket.""" @@ -1065,7 +1080,8 @@ def __del__(self) -> None: the object is lost due to leaving a function scope. It's always recommended that the user calls shutdown(). """ - self.shutdown() + if not sys.is_finalizing(): + self.shutdown() def invalidate_kv_cache(self) -> bool: """Invalidate reusable caches in vLLM (e.g., prefix/KV cache) after weight updates. diff --git a/nemo_rl/models/generation/vllm/vllm_worker.py b/nemo_rl/models/generation/vllm/vllm_worker.py index 3518338b671..7559c5c3e5a 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker.py +++ b/nemo_rl/models/generation/vllm/vllm_worker.py @@ -22,9 +22,8 @@ import threading import time import traceback -from collections import deque from concurrent.futures import Future, ThreadPoolExecutor -from typing import Any, Optional, cast +from typing import Any, Literal, Optional, cast import ray import torch @@ -58,18 +57,24 @@ from nemo_rl.models.policy.utils import is_vllm_v1_engine_enabled from nemo_rl.utils.nsys import wrap_with_nvtx_name from nemo_rl.utils.nvml import log_gpu_memory_diagnostics -from nemo_rl.utils.weight_transfer_s3_manifest import ( +from nemo_rl.utils.weight_transfer_remote_sparse import ( G_VLLM_REFIT_API_KEY_HEADER, G_VLLM_REFIT_FLUSH_PATH, G_VLLM_REFIT_S3_MANIFEST_PATH, + decode_sparse_payload, download_s3_refit_payload, merge_vllm_refit_receiver_timing, + refit_env_int, vllm_refit_api_key, ) - -G_REFIT_APPLY_QUEUE_DEPTH_ENV = "NRL_REFIT_APPLY_QUEUE_DEPTH" -G_REFIT_APPLY_BATCH_SIZE_ENV = "NRL_REFIT_APPLY_BATCH_SIZE" -G_REFIT_BATCH_STAGING_DIR_ENV = "NRL_REFIT_BATCH_STAGING_DIR" +from nemo_rl.utils.weight_transfer_zmq import ( + G_VLLM_REFIT_CHECKSUM_HEADER, + G_VLLM_REFIT_PAYLOAD_HEADER, + G_VLLM_REFIT_PRODUCER_HEADER, + G_VLLM_REFIT_TRANSFER_HEADER, + G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, + ZmqSparseRefitServer, +) logger = logging.getLogger(__name__) @@ -267,27 +272,24 @@ def __init__( if bundle_indices is not None and len(bundle_indices) == 1: bind_to_gpu_numa(int(ray.get_gpu_ids()[0])) - self._refit_apply_queue_lock = threading.Lock() + self._refit_apply_queue_condition = threading.Condition() self._refit_apply_executor = ThreadPoolExecutor(max_workers=1) - self._refit_apply_futures: deque[Future[dict[str, Any]]] = deque() + self._refit_apply_futures: list[Future[dict[str, Any]]] = [] self._refit_apply_pending_payloads: list[bytes] = [] - self._refit_apply_payload_count = 0 - self._refit_apply_batch_count = 0 + self._refit_seen_payloads: dict[tuple[str, int, int], str] = {} self._refit_workers_share_node = False - self._refit_apply_queue_depth = int( - os.getenv(G_REFIT_APPLY_QUEUE_DEPTH_ENV) or 2 + self._refit_apply_queue_depth = refit_env_int( + "NRL_REFIT_APPLY_QUEUE_DEPTH", default=2 + ) + self._refit_apply_batch_size = refit_env_int( + "NRL_REFIT_APPLY_BATCH_SIZE", default=8 ) - self._refit_apply_batch_size = int(os.getenv(G_REFIT_APPLY_BATCH_SIZE_ENV) or 8) self._refit_batch_staging_dir = ( - os.getenv(G_REFIT_BATCH_STAGING_DIR_ENV) or "/dev/shm" + os.getenv("NRL_REFIT_BATCH_STAGING_DIR") or "/dev/shm" ) - if self._refit_apply_queue_depth < 1: - raise ValueError(f"{G_REFIT_APPLY_QUEUE_DEPTH_ENV} must be >= 1.") - if self._refit_apply_batch_size < 1: - raise ValueError(f"{G_REFIT_APPLY_BATCH_SIZE_ENV} must be >= 1.") - self.refit_server_base_url: str | None = None - self.refit_server: Any | None = None - self.refit_server_thread: threading.Thread | None = None + self._refit_http_server: tuple[Any, threading.Thread, str] | None = None + self._zmq_refit_server: tuple[ZmqSparseRefitServer, str] | None = None + self._refit_async_loop: asyncio.AbstractEventLoop | None = None self._init_config( config, bundle_indices, fraction_of_gpus, seed, extra_env_vars @@ -698,43 +700,54 @@ def _get_raw_spec_counters(self) -> dict[str, float | list[float]]: def _enqueue_sparse_payload_apply( self, payload: bytes, + payload_key: tuple[str, int, int], + checksum: str, ) -> dict[str, Any]: completed: list[Future[dict[str, Any]]] = [] - with self._refit_apply_queue_lock: - while self._refit_apply_futures and ( - self._refit_apply_futures[0].done() - or len(self._refit_apply_futures) >= self._refit_apply_queue_depth + submitted = None + with self._refit_apply_queue_condition: + seen_checksum = self._refit_seen_payloads.get(payload_key) + if seen_checksum is not None: + if seen_checksum != checksum: + raise ValueError( + "A sparse refit payload ID was reused with different data." + ) + return {"ok": True, "payloads": 0, "duplicate": True} + while ( + len(self._refit_apply_futures) >= self._refit_apply_queue_depth + and not self._refit_apply_futures[0].done() ): - completed.append(self._refit_apply_futures.popleft()) + self._refit_apply_queue_condition.wait() + while self._refit_apply_futures and self._refit_apply_futures[0].done(): + completed.append(self._refit_apply_futures.pop(0)) response = self._collect_refit_apply_results(completed) + self._refit_seen_payloads[payload_key] = checksum self._refit_apply_pending_payloads.append(payload) - self._refit_apply_payload_count += 1 if len(self._refit_apply_pending_payloads) == self._refit_apply_batch_size: - self._submit_pending_sparse_payloads() + submitted = self._submit_pending_sparse_payloads() + if submitted is not None: + submitted.add_done_callback(self._notify_refit_apply_waiters) return response - def _submit_pending_sparse_payloads(self) -> None: + def _submit_pending_sparse_payloads(self) -> Future[dict[str, Any]]: payloads = tuple(self._refit_apply_pending_payloads) self._refit_apply_pending_payloads.clear() - self._refit_apply_futures.append( - self._refit_apply_executor.submit( - self.update_weights_from_serialized_sparse_payloads, - payloads, - False, - ) + future = self._refit_apply_executor.submit( + self.update_weights_from_serialized_sparse_payloads, + payloads, ) - self._refit_apply_batch_count += 1 + self._refit_apply_futures.append(future) + return future + + def _notify_refit_apply_waiters(self, _future: Future[dict[str, Any]]) -> None: + with self._refit_apply_queue_condition: + self._refit_apply_queue_condition.notify_all() def _collect_refit_apply_results( self, futures: list[Future[dict[str, Any]]], - *, - synchronize: bool = False, ) -> dict[str, Any]: results = [future.result() for future in futures] - if synchronize and results: - assert self.llm is not None - self.llm.collective_rpc("synchronize_device", args=()) timing: dict[str, float] = {} merge_vllm_refit_receiver_timing(timing, results, maximum=False) return { @@ -745,25 +758,105 @@ def _collect_refit_apply_results( @staticmethod def _refit_collective_response(worker_results: Any) -> dict[str, Any]: - return { + results = cast(list[dict[str, Any]], worker_results) + response = { "ok": True, - **merge_vllm_refit_receiver_timing( - {}, cast(list[dict[str, Any]], worker_results), maximum=True - ), + **merge_vllm_refit_receiver_timing({}, results, maximum=True), } + if any("verification_candidates" in result for result in results): + response["verification_candidates"] = max( + (int(result["verification_candidates"]) for result in results), + default=0, + ) + for key in ( + "verification_samples", + "verification_exact_mismatches", + "verification_mismatches", + "verification_abs_sum", + ): + response[key] = sum(result[key] for result in results) + response["verification_max_abs"] = max( + (float(result["verification_max_abs"]) for result in results), + default=0.0, + ) + return response + + def _refit_collective_rpc( + self, + method: str, + args: tuple[Any, ...], + ) -> Any: + return self.llm.collective_rpc(method, args=args) + + def update_weights_from_serialized_sparse_payloads( + self, + serialized_payloads: tuple[bytes, ...], + ) -> dict[str, Any]: + """Apply a FIFO batch of sparse deltas through one collective RPC.""" + if self.llm is None: + raise RuntimeError("vLLM is not initialized on this worker.") + if not self._refit_workers_share_node: + results = [ + self._refit_collective_response( + self._refit_collective_rpc( + "update_weights_from_serialized_sparse_payload", + (payload,), + ) + ) + for payload in serialized_payloads + ] + timing: dict[str, float] = {} + merge_vllm_refit_receiver_timing(timing, results, maximum=False) + return {"ok": True, "payloads": len(serialized_payloads), **timing} + + with tempfile.TemporaryDirectory( + prefix="nemo_rl_refit_", dir=self._refit_batch_staging_dir + ) as staging_dir: + paths = [] + for index, payload in enumerate(serialized_payloads): + path = os.path.join(staging_dir, str(index)) + with open(path, "wb") as staged: + staged.write(payload) + paths.append(path) + try: + response = self._refit_collective_response( + self._refit_collective_rpc( + "update_weights_from_sparse_payload_files", + tuple(paths), + ) + ) + except Exception: + # Drain peers before TemporaryDirectory removes shared batch files. + self._refit_collective_rpc("synchronize_device", ()) + raise + response["payloads"] = len(serialized_payloads) + return response def _flush_queued_sparse_payloads(self) -> dict[str, Any]: started = time.perf_counter() - with self._refit_apply_queue_lock: + submitted = None + with self._refit_apply_queue_condition: if self._refit_apply_pending_payloads: - self._submit_pending_sparse_payloads() + submitted = self._submit_pending_sparse_payloads() futures = list(self._refit_apply_futures) self._refit_apply_futures.clear() - payload_count = self._refit_apply_payload_count - batch_count = self._refit_apply_batch_count - self._refit_apply_payload_count = 0 - self._refit_apply_batch_count = 0 - response = self._collect_refit_apply_results(futures, synchronize=True) + self._refit_apply_queue_condition.notify_all() + payload_count = len(self._refit_seen_payloads) + batch_count = ( + payload_count + self._refit_apply_batch_size - 1 + ) // self._refit_apply_batch_size + if submitted is not None: + submitted.add_done_callback(self._notify_refit_apply_waiters) + response = self._collect_refit_apply_results(futures) + if futures: + assert self.llm is not None + response.update( + self._refit_collective_response( + self._refit_collective_rpc("finish_sparse_delta_refit", ()) + ) + ) + with self._refit_apply_queue_condition: + self._refit_seen_payloads.clear() response.update( payloads=payload_count, batches=batch_count, @@ -774,25 +867,20 @@ def _flush_queued_sparse_payloads(self) -> dict[str, Any]: "REFIT_RECEIVER_TIMING " f"payloads={payload_count} batches={batch_count} " f"total_s={response['seconds']:.3f} " - f"payload_total_s={response.get('receiver_total_s', 0.0):.3f}", + f"payload_total_s={response.get('receiver_total_s', 0.0):.3f} " + f"delta_verify_candidates=" + f"{response.get('verification_candidates', 0)} " + f"delta_verify_samples={response.get('verification_samples', 0)} " + f"delta_verify_exact_mismatches=" + f"{response.get('verification_exact_mismatches', 0)} " + f"delta_verify_mismatches=" + f"{response.get('verification_mismatches', 0)} " + f"delta_verify_max_abs=" + f"{response.get('verification_max_abs', 0.0):.8g}", flush=True, ) return response - def update_weights_from_serialized_sparse_payload( - self, serialized_payload: bytes, synchronize: bool = True - ) -> dict[str, Any]: - return self.update_weights_from_serialized_sparse_payloads( - (serialized_payload,), synchronize - ) - - def update_weights_from_serialized_sparse_payloads( - self, - serialized_payloads: tuple[bytes, ...], - synchronize: bool = True, - ) -> dict[str, Any]: - raise NotImplementedError - async def _apply_s3_manifest_payload( self, manifest: dict[str, Any], @@ -800,8 +888,40 @@ async def _apply_s3_manifest_payload( started = time.perf_counter() body = await asyncio.to_thread(download_s3_refit_payload, manifest) download_s = time.perf_counter() - started - result = await asyncio.to_thread(self._enqueue_sparse_payload_apply, body) - result.update(payloads=1, receiver_s3_download_s=download_s) + key = str(manifest["key"]) + checksum = str(manifest["checksum"]) + result = await asyncio.to_thread( + self._enqueue_sparse_payload_apply, + body, + (key, -1, -1), + checksum, + ) + result["receiver_s3_download_s"] = download_s + return result + + async def _apply_zmq_payload(self, raw_request: Any) -> dict[str, Any]: + headers = raw_request.headers + transfer_id = headers.get(G_VLLM_REFIT_TRANSFER_HEADER, "") + producer_id = int(headers.get(G_VLLM_REFIT_PRODUCER_HEADER, "-1")) + payload_id = int(headers.get(G_VLLM_REFIT_PAYLOAD_HEADER, "-1")) + checksum = headers.get(G_VLLM_REFIT_CHECKSUM_HEADER, "") + if not transfer_id or producer_id < 0 or payload_id < 0 or not checksum: + raise ValueError("Missing or invalid ZeroMQ sparse refit payload headers.") + compressed = await raw_request.body() + started = time.perf_counter() + payload = await asyncio.to_thread( + decode_sparse_payload, + compressed, + checksum, + ) + decode_s = time.perf_counter() - started + result = await asyncio.to_thread( + self._enqueue_sparse_payload_apply, + payload, + (transfer_id, producer_id, payload_id), + checksum, + ) + result["receiver_zmq_decode_s"] = decode_s return result def _setup_vllm_refit_api_server(self, app: Any) -> None: @@ -812,7 +932,12 @@ def _setup_vllm_refit_api_server(self, app: Any) -> None: self.cfg["vllm_cfg"].get("http_refit_api_key_env_var") ) - async def respond(raw_request: Request, *, flush: bool = False) -> JSONResponse: + async def respond( + raw_request: Request, + action: Literal["s3", "flush", "zmq"], + ) -> JSONResponse: + if self.cfg["vllm_cfg"]["async_engine"]: + self._refit_async_loop = asyncio.get_running_loop() if ( token is not None and raw_request.headers.get(G_VLLM_REFIT_API_KEY_HEADER) != token @@ -821,11 +946,14 @@ async def respond(raw_request: Request, *, flush: bool = False) -> JSONResponse: content={"ok": False, "error": "unauthorized"}, status_code=403 ) try: - result = ( - await asyncio.to_thread(self._flush_queued_sparse_payloads) - if flush - else await self._apply_s3_manifest_payload(await raw_request.json()) - ) + if action == "s3": + result = await self._apply_s3_manifest_payload( + await raw_request.json() + ) + elif action == "zmq": + result = await self._apply_zmq_payload(raw_request) + else: + result = await asyncio.to_thread(self._flush_queued_sparse_payloads) except Exception as exc: result = {"ok": False, "error": str(exc)} return JSONResponse( @@ -835,14 +963,44 @@ async def respond(raw_request: Request, *, flush: bool = False) -> JSONResponse: @app.post(G_VLLM_REFIT_S3_MANIFEST_PATH) async def apply_s3_manifest_refit(raw_request: Request) -> JSONResponse: - return await respond(raw_request) + return await respond(raw_request, "s3") @app.post(G_VLLM_REFIT_FLUSH_PATH) async def flush_sparse_delta_refit(raw_request: Request) -> JSONResponse: - return await respond(raw_request, flush=True) + return await respond(raw_request, "flush") + + @app.post(G_VLLM_REFIT_ZMQ_PAYLOAD_PATH) + async def apply_zmq_sparse_refit(raw_request: Request) -> JSONResponse: + return await respond(raw_request, "zmq") def report_refit_server_base_url(self) -> str | None: - return self.refit_server_base_url + return self._refit_http_server[2] if self._refit_http_server else None + + def start_zmq_sparse_refit_relay(self, refit_urls: list[str]) -> str: + if self._zmq_refit_server is not None: + return self._zmq_refit_server[1] + port = self.cfg["vllm_cfg"].get( + "zmq_refit_server_port" + ) or _get_free_port_local( + self.cfg.get("port_range_low", DEFAULT_GENERATION_PORT_RANGE_LOW), + self.cfg.get("port_range_high", DEFAULT_GENERATION_PORT_RANGE_HIGH), + ) + server = ZmqSparseRefitServer( + refit_urls, + bind_address=f"tcp://0.0.0.0:{port}", + api_key_env_var=self.cfg["vllm_cfg"].get("http_refit_api_key_env_var"), + timeout_s=float(os.getenv("NRL_REFIT_ZMQ_TIMEOUT_S") or 600.0), + ) + server.start() + address = f"tcp://{_get_node_ip_local()}:{port}" + self._zmq_refit_server = (server, address) + print(f"Starting vLLM ZeroMQ refit relay on {address}", flush=True) + return address + + def stop_zmq_sparse_refit_relay(self) -> None: + if self._zmq_refit_server is not None: + self._zmq_refit_server[0].close() + self._zmq_refit_server = None class VllmGenerationWorkerImpl(BaseVllmGenerationWorker): @@ -859,7 +1017,7 @@ def post_init(self): self.llm.collective_rpc( "load_mtp_weights_from_disk", args=(self.model_name,) ) - if self.cfg["vllm_cfg"].get("expose_http_refit_server"): + if self.cfg.get("refit_transport") is not None: self._refit_workers_share_node = ( len(set(self.llm.collective_rpc("report_node_hostname", args=()))) == 1 ) @@ -887,10 +1045,9 @@ def _setup_vllm_refit_server(self) -> None: ) thread = threading.Thread(target=server.run, daemon=True) thread.start() - self.refit_server_base_url = f"http://{_get_node_ip_local()}:{port}" - self.refit_server = server - self.refit_server_thread = thread - print(f"Starting vLLM refit server on {self.refit_server_base_url}", flush=True) + base_url = f"http://{_get_node_ip_local()}:{port}" + self._refit_http_server = (server, thread, base_url) + print(f"Starting vLLM refit server on {base_url}", flush=True) def init_collective( self, @@ -1245,57 +1402,6 @@ def update_weights_from_collective(self) -> bool: traceback.print_exc() return False - def update_weights_from_serialized_sparse_payloads( - self, - serialized_payloads: tuple[bytes, ...], - synchronize: bool = True, - ) -> dict[str, Any]: - """Apply a FIFO batch of S3 sparse deltas through one collective RPC.""" - if self.llm is None: - raise RuntimeError( - "Attempting to update weights with either an uninitialized vLLM " - "or non-model-owner" - ) - if not serialized_payloads: - raise ValueError("A sparse refit batch must contain at least one payload.") - if not self._refit_workers_share_node: - results = [ - self._refit_collective_response( - self.llm.collective_rpc( - "update_weights_from_serialized_sparse_payload", - args=(payload, False), - ) - ) - for payload in serialized_payloads - ] - if synchronize: - self.llm.collective_rpc("synchronize_device", args=()) - timing: dict[str, float] = {} - merge_vllm_refit_receiver_timing(timing, results, maximum=False) - return {"ok": True, "payloads": len(serialized_payloads), **timing} - - paths: list[str] = [] - try: - for payload in serialized_payloads: - fd, path = tempfile.mkstemp( - prefix="nemo_rl_refit_", dir=self._refit_batch_staging_dir - ) - paths.append(path) - with os.fdopen(fd, "wb") as staged: - staged.write(payload) - response = self._refit_collective_response( - self.llm.collective_rpc( - "update_weights_from_sparse_payload_files", - args=tuple(paths), - kwargs={"synchronize": synchronize}, - ) - ) - finally: - for path in paths: - os.unlink(path) - response["payloads"] = len(serialized_payloads) - return response - def reset_prefix_cache(self): """Reset the prefix cache of vLLM engine.""" assert self.llm is not None, ( @@ -1361,14 +1467,16 @@ def wake_up(self, **kwargs): def shutdown(self) -> bool: """Clean up vLLM resources.""" try: - if self.refit_server is not None: - self.refit_server.should_exit = True + self.stop_zmq_sparse_refit_relay() + if self._refit_http_server is not None: + self._refit_http_server[0].should_exit = True self._flush_queued_sparse_payloads() self._refit_apply_executor.shutdown(wait=True) - if self.refit_server_thread is not None: - self.refit_server_thread.join(timeout=5.0) + if self._refit_http_server is not None: + self._refit_http_server[1].join(timeout=5.0) + self._refit_http_server = None if self.llm is not None: # Clean up extension resources (e.g., ZMQ sockets) diff --git a/nemo_rl/models/generation/vllm/vllm_worker_async.py b/nemo_rl/models/generation/vllm/vllm_worker_async.py index eae64f3f079..f9baf78fc4d 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker_async.py +++ b/nemo_rl/models/generation/vllm/vllm_worker_async.py @@ -218,6 +218,17 @@ def __init__( self.llm = None self.vllm_device_ids = None + def _refit_collective_rpc( + self, + method: str, + args: tuple[Any, ...], + ) -> Any: + if self._refit_async_loop is None: + raise RuntimeError("The async vLLM refit server loop is not initialized.") + return asyncio.run_coroutine_threadsafe( + self.llm.collective_rpc(method, args=args), self._refit_async_loop + ).result() + def _return_routed_experts_enabled(self) -> bool: engine_args = getattr(self, "llm_async_engine_args", None) if bool(getattr(engine_args, "enable_return_routed_experts", False)): @@ -429,6 +440,9 @@ async def post_init_async(self): await self.llm.collective_rpc( "load_mtp_weights_from_disk", args=(self.model_name,) ) + if self.cfg.get("refit_transport") is not None: + hostnames = await self.llm.collective_rpc("report_node_hostname", args=()) + self._refit_workers_share_node = len(set(hostnames)) == 1 async def get_reserved_url(self) -> Optional[str]: """Return the URL from the reserved socket, available before model loading.""" @@ -439,6 +453,11 @@ async def get_reserved_url(self) -> Optional[str]: async def report_dp_openai_server_base_url(self) -> Optional[str]: return self.base_url + def report_refit_server_base_url(self) -> str | None: + if self.cfg.get("refit_transport") is None or self.base_url is None: + return None + return self.base_url.removesuffix("/v1") + # ruff: noqa def _setup_vllm_openai_api_server(self, app: FastAPI) -> FastAPI: from copy import deepcopy @@ -911,6 +930,8 @@ def _setup_vllm_server(self) -> "tuple[threading.Thread, str, uvicorn.Server]": app = FastAPI() app = self._setup_vllm_openai_api_server(app) + if self.cfg.get("refit_transport") is not None: + self._setup_vllm_refit_api_server(app) ######################################## # Server spinup @@ -1535,6 +1556,10 @@ async def wake_up_async(self, **kwargs): async def shutdown(self) -> bool: """Clean up vLLM resources.""" try: + self.stop_zmq_sparse_refit_relay() + await asyncio.to_thread(self._flush_queued_sparse_payloads) + self._refit_apply_executor.shutdown(wait=True) + if self.llm is not None: # Clean up extension resources (e.g., ZMQ sockets) await self.llm.collective_rpc("cleanup", args=tuple()) diff --git a/nemo_rl/models/policy/lm_policy.py b/nemo_rl/models/policy/lm_policy.py index c1f1a90084b..226641b2039 100644 --- a/nemo_rl/models/policy/lm_policy.py +++ b/nemo_rl/models/policy/lm_policy.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. import os +import sys import warnings from collections import defaultdict from contextlib import nullcontext @@ -940,10 +941,11 @@ 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 init_remote_sparse_delta_baseline(self) -> list[ray.ObjectRef]: - """Initialize source-side sparse-delta baselines for remote S3 refit.""" - return self._run_s3_refit_workers( + def init_remote_sparse_delta_baseline(self, transport: str) -> list[ray.ObjectRef]: + """Initialize source-side sparse-delta baselines for remote refit.""" + return self._run_remote_sparse_refit_workers( "init_remote_sparse_delta_baseline", + transport=transport, ) def finish_inference(self) -> None: @@ -1069,17 +1071,21 @@ def set_rollout_num_gpus_per_engine(self, num_gpus_per_engine: int) -> None: ) ) - def stream_sparse_weights_via_s3_manifest( + def stream_remote_sparse_weights( self, - refit_urls: list[str], + transport: str, + targets: list[str], *, - api_key_env_var: Optional[str] = None, - timeout_s: float = 600.0, + transfer_id: str, + api_key_env_var: Optional[str], + timeout_s: float, ) -> list[ray.ObjectRef]: - """Upload vLLM refit payloads to S3 and post receiver manifests.""" - return self._run_s3_refit_workers( - "stream_sparse_weights_via_s3_manifest", - refit_urls=refit_urls, + """Stream sparse deltas through the selected remote value plane.""" + return self._run_remote_sparse_refit_workers( + "stream_remote_sparse_weights", + transport=transport, + targets=targets, + transfer_id=transfer_id, api_key_env_var=api_key_env_var, timeout_s=timeout_s, ) @@ -1089,7 +1095,7 @@ def finish_remote_sparse_delta_sync(self, succeeded: bool) -> list[ray.ObjectRef "finish_remote_sparse_delta_sync", succeeded=succeeded ) - def _run_s3_refit_workers( + def _run_remote_sparse_refit_workers( self, method_name: str, **common_kwargs: Any, @@ -1193,7 +1199,7 @@ def __del__(self) -> None: the object is lost due to leaving a function scope. It's always recommended that the user calls worker_group.shutdown(). """ - if hasattr(self, "worker_group"): + if not sys.is_finalizing() and hasattr(self, "worker_group"): self.worker_group.shutdown(cleanup_method="shutdown") def start_gpu_profiling(self) -> None: diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index a1aaabd8af7..7cef71bf17d 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -64,7 +64,6 @@ ) from nemo_rl.models.megatron.pipeline_parallel import ( broadcast_loss_metrics_from_last_stage, - broadcast_obj_from_pp_rank, broadcast_tensors_from_last_stage, ) from nemo_rl.models.megatron.router_replay import router_replay_enabled @@ -99,11 +98,13 @@ from nemo_rl.utils.packed_tensor import packed_broadcast_producer from nemo_rl.utils.r3_trace import maybe_r3_trace_stage from nemo_rl.utils.timer import Timer -from nemo_rl.utils.weight_transfer_s3_manifest import ( +from nemo_rl.utils.weight_transfer_remote_sparse import ( + SparseDeltaStreamResult, init_sparse_delta_baseline_from_iterator, stream_sparse_delta_payloads_via_s3_manifest, ) from nemo_rl.utils.weight_transfer_sparse_codec import DeltaCompressionTracker +from nemo_rl.utils.weight_transfer_zmq import stream_sparse_delta_payloads_via_zmq TokenizerType = TypeVar("TokenizerType", bound=PreTrainedTokenizerBase) @@ -369,11 +370,12 @@ def __init__( self.is_generation_colocated = runtime_config.is_generation_colocated self.final_padded_vocab_size = runtime_config.final_padded_vocab_size self.sampling_params = runtime_config.sampling_params - delta_config = self.cfg.get("generation", {}).get("delta_compression") + generation_config = self.cfg.get("generation") + delta_config = None + if generation_config is not None and generation_config["backend"] == "vllm": + delta_config = cast(VllmConfig, generation_config).get("delta_compression") self.delta_weight_transfer_tracker = ( - DeltaCompressionTracker(delta_config) - if delta_config and delta_config["enabled"] - else None + DeltaCompressionTracker(delta_config) if delta_config else None ) self.defer_fp32_logits = self.cfg["megatron_cfg"].get( @@ -1811,38 +1813,67 @@ def _get_model_config(self): def init_remote_sparse_delta_baseline( self, *, - shard_rank: int = 0, - shard_count: int = 1, + shard_rank: int, + shard_count: int, + transport: str, ) -> None: - """Initialize the source-side baseline for remote sparse S3 refit.""" + """Initialize the source-side baseline for remote sparse refit.""" + tracker = self.delta_weight_transfer_tracker + assert tracker is not None init_sparse_delta_baseline_from_iterator( self._iter_params_with_optional_kv_scales(), - delta_tracker=self.delta_weight_transfer_tracker, + delta_tracker=tracker, shard_rank=shard_rank, shard_count=shard_count, + transport=transport, ) @torch.no_grad() - @wrap_with_nvtx_name("megatron_policy_worker/stream_sparse_weights_via_s3_manifest") - def stream_sparse_weights_via_s3_manifest( + @wrap_with_nvtx_name("megatron_policy_worker/stream_remote_sparse_weights") + def stream_remote_sparse_weights( self, - refit_urls: list[str], + transport: str, + targets: list[str], *, - api_key_env_var: Optional[str] = None, - timeout_s: float = 600.0, - shard_rank: int = 0, - shard_count: int = 1, - ) -> dict[str, Any]: - """Upload vLLM refit payloads to S3 and post receiver manifests.""" - return stream_sparse_delta_payloads_via_s3_manifest( + transfer_id: str, + api_key_env_var: Optional[str], + timeout_s: float, + shard_rank: int, + shard_count: int, + ) -> SparseDeltaStreamResult: + """Stream compressed sparse deltas through the selected value plane.""" + tracker = self.delta_weight_transfer_tracker + assert tracker is not None + streamer = { + "s3": stream_sparse_delta_payloads_via_s3_manifest, + "zmq": stream_sparse_delta_payloads_via_zmq, + }.get(transport) + if streamer is None: + raise ValueError( + f"Unsupported remote sparse refit transport {transport!r}." + ) + result = streamer( self._iter_params_with_optional_kv_scales(), - delta_tracker=self.delta_weight_transfer_tracker, - refit_urls=refit_urls, + delta_tracker=tracker, + refit_targets=targets, + transfer_id=transfer_id, api_key_env_var=api_key_env_var, timeout_s=timeout_s, shard_rank=shard_rank, shard_count=shard_count, ) + # HF export can enqueue work on auxiliary streams. Drain the device before + # this actor returns so every policy rank resumes training from the same + # completed CUDA boundary. + if torch.cuda.is_available(): + torch.cuda.synchronize() + return result + + def _get_refit_conversion_tasks(self) -> list[Any]: + if self.refit_conversion_tasks is None: + tasks = self.megatron_bridge.get_conversion_tasks([self.model]) + self.refit_conversion_tasks = [task for task in tasks if task is not None] + return self.refit_conversion_tasks def finish_remote_sparse_delta_sync(self, *, succeeded: bool) -> None: tracker = self.delta_weight_transfer_tracker @@ -1867,14 +1898,10 @@ def _calculate_refit_param_info(self) -> list[tuple[str, int]]: Returns: List of (parameter_name, size_in_bytes) tuples. """ - self.refit_conversion_tasks = [ - task - for task in self.megatron_bridge.get_conversion_tasks([self.model]) - if task is not None - ] + conversion_tasks = self._get_refit_conversion_tasks() param_info = [] - def calculate_size_in_bytes(param, tp_size, ep_size): + def calculate_size_in_bytes(param, mapping): if param is None: # need to broadcast for other pp ranks size_in_bytes = None @@ -1889,20 +1916,23 @@ def calculate_size_in_bytes(param, tp_size, ep_size): } scale = prec_to_bytes[self.dtype] / prec_to_bytes[param.dtype] size_in_bytes = ( - param.element_size() * param.numel() * tp_size * ep_size * scale + param.element_size() + * param.numel() + * mapping.tp_size + * (mapping.ep_size if mapping.is_expert else 1) + * scale ) - # Broadcast size_in_bytes across pipeline parallel ranks - return broadcast_obj_from_pp_rank(size_in_bytes) + # Match Megatron Bridge export semantics for tied or replicated weights. + return mapping.broadcast_obj_from_pp_rank(size_in_bytes) - for task in self.refit_conversion_tasks: + for task in conversion_tasks: param_info.append( ( task.param_name, calculate_size_in_bytes( task.param_weight, - task.mapping.tp_size, - task.mapping.ep_size if task.mapping.is_expert else 1, + task.mapping, ), ) ) @@ -1924,7 +1954,7 @@ def _iter_params_with_optional_kv_scales( base_iter = self.megatron_bridge.export_hf_weights( [self.model], show_progress=False, - conversion_tasks=self.refit_conversion_tasks, # used for metadata caching + conversion_tasks=self._get_refit_conversion_tasks(), ) # Yield the original parameters first. diff --git a/nemo_rl/utils/weight_transfer_remote_sparse.py b/nemo_rl/utils/weight_transfer_remote_sparse.py new file mode 100644 index 00000000000..f839c1d8dcf --- /dev/null +++ b/nemo_rl/utils/weight_transfer_remote_sparse.py @@ -0,0 +1,630 @@ +# 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. + +"""Shared sparse payload pipeline and control plane for remote vLLM refit.""" + +import hashlib +import io +import os +import threading +import time +from collections.abc import Iterable, Iterator, Mapping, Sequence +from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait +from contextlib import suppress +from functools import cache +from typing import Any, Callable, TypedDict +from urllib.parse import quote + +import requests +import torch +import zstandard +from urllib3.util.retry import Retry + +from nemo_rl.utils.packed_tensor import get_target_packed_tensor_size +from nemo_rl.utils.weight_transfer_sparse_codec import ( + DeltaCompressionTracker, + NamedTensor, + TensorBatch, +) + +G_VLLM_REFIT_S3_MANIFEST_PATH = "/nemo-rl/refit/s3-manifest" +G_VLLM_REFIT_FLUSH_PATH = "/nemo-rl/refit/flush" +G_VLLM_REFIT_API_KEY_HEADER = "x-nemo-rl-refit-key" +_CONTROL_SESSION_LOCAL = threading.local() +_S3_PART_SIZE = 64 * 1024**2 +_S3_MEMORY_LIMIT = 2 * 1024**3 + + +class SparseDeltaStreamResult(TypedDict): + payloads: int + changed_elements: int + total_elements: int + + +@cache +def _s3_client(region: str) -> Any: + from awscrt.auth import AwsCredentialsProvider + from awscrt.io import ClientBootstrap, DefaultHostResolver, EventLoopGroup + from awscrt.s3 import S3Client, create_default_s3_signing_config + + event_loop_group = EventLoopGroup() + bootstrap = ClientBootstrap( + event_loop_group, + DefaultHostResolver(event_loop_group), + ) + return S3Client( + bootstrap=bootstrap, + region=region, + signing_config=create_default_s3_signing_config( + region=region, + credential_provider=AwsCredentialsProvider.new_default_chain(bootstrap), + ), + part_size=_S3_PART_SIZE, + multipart_upload_threshold=_S3_PART_SIZE, + throughput_target_gbps=10.0, + memory_limit=_S3_MEMORY_LIMIT, + ) + + +class _S3ObjectStore: + def __init__(self, *, bucket: str, region: str) -> None: + from awscrt.s3 import S3RequestType + + self.bucket = bucket + self.region = region + self._client = _s3_client(region) + self._request_type = S3RequestType + + def put(self, key: str, body: bytes) -> None: + self._client.make_request( + type=self._request_type.PUT_OBJECT, + request=self._request("PUT", key, body), + ).finished_future.result() + + def get(self, key: str) -> bytes: + from awscrt.http import HttpHeaders + + body = bytearray() + + def on_headers( + status_code: int, + headers: list[tuple[str, str]], + **_kwargs: Any, + ) -> None: + nonlocal body + if status_code != 200: + raise RuntimeError(f"S3 GET returned HTTP {status_code}.") + length = HttpHeaders(headers).get("content-length") + if length is None: + raise RuntimeError("S3 GET response omitted content-length.") + body = bytearray(int(length)) + + def on_body(chunk: bytes, offset: int, **_kwargs: Any) -> None: + body[offset : offset + len(chunk)] = chunk + + self._client.make_request( + type=self._request_type.GET_OBJECT, + request=self._request("GET", key), + on_headers=on_headers, + on_body=on_body, + ).finished_future.result() + return bytes(body) + + def delete(self, key: str) -> None: + self._client.make_request( + type=self._request_type.DEFAULT, + request=self._request("DELETE", key), + operation_name="DeleteObject", + ).finished_future.result() + + def _request(self, method: str, key: str, body: bytes | None = None) -> Any: + from awscrt.http import HttpHeaders, HttpRequest + + headers = HttpHeaders( + [("host", f"{self.bucket}.s3.{self.region}.amazonaws.com")] + ) + if body is not None: + headers.add("content-length", str(len(body))) + headers.add("content-type", "application/octet-stream") + elif method == "DELETE": + headers.add("content-length", "0") + return HttpRequest( + method, + f"/{quote(key, safe='/~')}", + headers, + io.BytesIO(body) if body is not None else None, + ) + + +def refit_env_int(name: str, *, default: int, min_value: int = 1) -> int: + value = int(os.getenv(name) or default) + if value < min_value: + raise ValueError(f"{name} must be >= {min_value}.") + return value + + +def sparse_payload_checksum(body: bytes) -> str: + return hashlib.blake2b(body, digest_size=16).hexdigest() + + +def decode_sparse_payload(body: bytes, checksum: str) -> bytes: + actual = sparse_payload_checksum(body) + if actual != checksum: + raise ValueError( + f"Sparse refit payload checksum mismatch: expected={checksum}, actual={actual}." + ) + decompressor = getattr(_CONTROL_SESSION_LOCAL, "zstd_decompressor", None) + if decompressor is None: + decompressor = zstandard.ZstdDecompressor() + _CONTROL_SESSION_LOCAL.zstd_decompressor = decompressor + return decompressor.decompress(body) + + +def iter_sparse_weight_chunks( + tensors: Iterable[NamedTensor], target_bytes: int +) -> Iterator[tuple[TensorBatch, float]]: + iterator = iter(tensors) + pending = None + while True: + started = time.perf_counter() + chunk = [pending] if pending is not None else [] + size = pending[1].numel() * pending[1].element_size() if pending else 0 + pending = None + for item in iterator: + item_size = item[1].numel() * item[1].element_size() + if chunk and size + item_size > target_bytes: + pending = item + break + chunk.append(item) + size += item_size + if size >= target_bytes: + break + export_pull_s = time.perf_counter() - started + if not chunk: + return + yield chunk, export_pull_s + + +def refit_http_session() -> requests.Session: + session = getattr(_CONTROL_SESSION_LOCAL, "session", None) + if session is None: + session = requests.Session() + adapter = requests.adapters.HTTPAdapter( + pool_connections=64, + pool_maxsize=64, + max_retries=Retry( + total=3, + backoff_factor=0.25, + status_forcelist=(502, 503, 504), + allowed_methods={"POST"}, + ), + ) + session.mount("http://", adapter) + session.mount("https://", adapter) + _CONTROL_SESSION_LOCAL.session = session + return session + + +@cache +def _get_manifest_s3_store(bucket: str, region: str) -> _S3ObjectStore: + return _S3ObjectStore(bucket=bucket, region=region) + + +def vllm_refit_api_key(api_key_env_var: str | None) -> str | None: + if not api_key_env_var: + return None + token = os.environ.get(api_key_env_var) + if not token: + raise RuntimeError( + "vLLM sparse refit API key env var " + f"{api_key_env_var!r} is configured but unset or empty." + ) + return token + + +def sparse_export_chunk_size( + delta_tracker: DeltaCompressionTracker, + transport: str, +) -> int: + requested = refit_env_int( + f"NRL_REFIT_{transport.upper()}_EXPORT_CHUNK_BYTES", + default=(1024 if transport == "zmq" else 256) * 1024**2, + min_value=1, + ) + if torch.cuda.is_available(): + requested = min(requested, get_target_packed_tensor_size()) + return min(requested, delta_tracker.sparse_bucket_size_bytes) + + +@cache +def _executor(key: str, workers: int) -> ThreadPoolExecutor: + return ThreadPoolExecutor(max_workers=workers, thread_name_prefix=f"nrl-{key}") + + +def init_sparse_delta_baseline_from_iterator( + iterator: Iterable[NamedTensor], + *, + delta_tracker: DeltaCompressionTracker, + shard_rank: int, + shard_count: int, + transport: str, +) -> None: + start_s = time.perf_counter() + export_chunk_size = sparse_export_chunk_size(delta_tracker, transport) + + chunk_count = 0 + export_pull_s = snapshot_s = 0.0 + for chunk_index, (chunk, pull_s) in enumerate( + iter_sparse_weight_chunks(iterator, export_chunk_size) + ): + chunk_count = chunk_index + 1 + export_pull_s += pull_s + if chunk_index % shard_count != shard_rank: + continue + started = time.perf_counter() + delta_tracker.snapshot_baseline(chunk) + snapshot_s += time.perf_counter() - started + print( + "REFIT_BASELINE_INIT " + f"event=end chunks={chunk_count} export_pull_s={export_pull_s:.3f} " + f"snapshot_s={snapshot_s:.3f} " + f"seconds={time.perf_counter() - start_s:.3f}", + flush=True, + ) + + +def stream_sparse_delta_payloads( + iterator: Iterable[NamedTensor], + *, + delta_tracker: DeltaCompressionTracker, + transport: str, + send_payload: Callable[[bytes, int], dict[str, Any]], + transfer_workers: int, + shard_rank: int, + shard_count: int, +) -> SparseDeltaStreamResult: + prefix = transport.upper() + encode_workers = refit_env_int( + f"NRL_REFIT_{prefix}_ENCODE_WORKERS", + default=max(2, min(8, os.cpu_count() or 8)), + ) + pipeline_workers = max(encode_workers, transfer_workers) + encode_executor = _executor(f"refit-{transport}-encode", encode_workers) + transfer_executor = _executor(f"refit-{transport}-transfer", transfer_workers) + export_chunk_size = sparse_export_chunk_size(delta_tracker, transport) + + def encode_chunk( + chunk: TensorBatch, + ) -> tuple[bytes | None, dict[str, float], int, int]: + started = time.perf_counter() + payload, changed_elements, total_elements = ( + delta_tracker.prepare_sparse_delta_payload(chunk) + ) + encode_s = time.perf_counter() - started + if not payload[2]: + return None, {"encode_s": encode_s}, changed_elements, total_elements + started = time.perf_counter() + buffer = io.BytesIO() + torch.save(payload, buffer) + raw_body = buffer.getvalue() + serialize_s = time.perf_counter() - started + started = time.perf_counter() + body = zstd_compress(raw_body, f"NRL_REFIT_{prefix}_ZSTD_THREADS") + compress_s = time.perf_counter() - started + return ( + body, + { + "encode_s": encode_s, + "serialize_s": serialize_s, + "compress_s": compress_s, + }, + changed_elements, + total_elements, + ) + + def transfer_payload( + encoded: tuple[bytes, dict[str, float]], payload_index: int + ) -> dict[str, Any]: + body, encode_timing = encoded + result = send_payload(body, payload_index) + result.update( + body_size=len(body), + **encode_timing, + ) + return result + + timing: dict[str, float] = {} + receiver_timing: dict[str, float] = {} + counts = { + "payloads": 0, + "wire_bytes": 0, + "changed_elements": 0, + "total_elements": 0, + } + chunk_count = 0 + export_pull_s = 0.0 + encode_inflight: dict[Any, int] = {} + transfer_inflight: set[Any] = set() + worker_errors: list[Exception] = [] + max_encode_inflight = encode_workers * 2 + + def collect_transfers(*, block: bool) -> None: + if not transfer_inflight: + return + if block: + completed, _ = wait(transfer_inflight, return_when=FIRST_COMPLETED) + else: + completed = {future for future in transfer_inflight if future.done()} + for future in completed: + transfer_inflight.remove(future) + try: + result = future.result() + except Exception as error: + worker_errors.append(error) + continue + counts["payloads"] += 1 + counts["wire_bytes"] += int(result["body_size"]) + for key, value in result.items(): + if key.endswith("_s"): + timing[key] = timing.get(key, 0.0) + float(value) + merge_vllm_refit_receiver_timing( + receiver_timing, [result["receiver"]], maximum=False + ) + + def drain_encodes() -> None: + completed, _ = wait(encode_inflight, return_when=FIRST_COMPLETED) + for future in completed: + payload_index = encode_inflight.pop(future) + try: + encoded = future.result() + except Exception as error: + worker_errors.append(error) + continue + body, encode_timing, changed_elements, total_elements = encoded + counts["changed_elements"] += changed_elements + counts["total_elements"] += total_elements + if body is not None: + transfer_inflight.add( + transfer_executor.submit( + transfer_payload, + (body, encode_timing), + payload_index, + ) + ) + collect_transfers(block=False) + + payload_index = 0 + stream_start = time.perf_counter() + try: + for chunk_index, (chunk, pull_s) in enumerate( + iter_sparse_weight_chunks(iterator, export_chunk_size) + ): + chunk_count = chunk_index + 1 + export_pull_s += pull_s + if chunk_index % shard_count != shard_rank: + continue + if len(encode_inflight) >= max_encode_inflight: + drain_encodes() + encode_inflight[encode_executor.submit(encode_chunk, chunk)] = payload_index + payload_index += 1 + + while encode_inflight: + drain_encodes() + while transfer_inflight: + collect_transfers(block=True) + if worker_errors: + raise worker_errors[0] + except Exception: + for future in (*encode_inflight, *transfer_inflight): + future.cancel() + if encode_inflight: + wait(encode_inflight) + if transfer_inflight: + wait(transfer_inflight) + raise + + timing = { + "total_s": time.perf_counter() - stream_start, + "export_pull_s": export_pull_s, + **timing, + "payloads": counts["payloads"], + "chunks": chunk_count, + "wire_mb": counts["wire_bytes"] / 1e6, + "pipeline_workers": pipeline_workers, + "encode_workers": encode_workers, + "export_chunk_mb": export_chunk_size / 1e6, + "shard_rank": shard_rank, + "shard_count": shard_count, + "changed_elements": counts["changed_elements"], + "total_elements": counts["total_elements"], + "changed_pct": 100.0 + * counts["changed_elements"] + / max(counts["total_elements"], 1), + } + timing.update(receiver_timing) + print( + f"REFIT_{prefix}_TIMING " + + " ".join(f"{key}={value}" for key, value in timing.items()), + flush=True, + ) + return { + "payloads": counts["payloads"], + "changed_elements": counts["changed_elements"], + "total_elements": counts["total_elements"], + } + + +def stream_sparse_delta_payloads_via_s3_manifest( + iterator: Iterable[NamedTensor], + *, + delta_tracker: DeltaCompressionTracker, + refit_targets: Sequence[str], + transfer_id: str, + api_key_env_var: str | None, + timeout_s: float, + shard_rank: int, + shard_count: int, +) -> SparseDeltaStreamResult: + urls = [url.strip().rstrip("/") for url in refit_targets if url.strip()] + if not urls: + raise ValueError("At least one vLLM S3 refit URL is required.") + bucket = os.getenv("NRL_REFIT_S3_BUCKET", "").strip() + if not bucket: + raise RuntimeError("NRL_REFIT_S3_BUCKET must be set for S3 refit.") + store = _get_manifest_s3_store( + bucket, + os.getenv("NRL_REFIT_S3_REGION", "us-east-1").strip() or "us-east-1", + ) + endpoint_urls = [f"{url}{G_VLLM_REFIT_S3_MANIFEST_PATH}" for url in urls] + object_prefix = os.getenv("NRL_REFIT_S3_PREFIX", "nemo-rl-refit").strip("/") + run_prefix = ( + f"{object_prefix}/{transfer_id}/{shard_rank:06d}" + if object_prefix + else f"{transfer_id}/{shard_rank:06d}" + ) + api_key = vllm_refit_api_key(api_key_env_var) + + def send_payload(body: bytes, payload_index: int) -> dict[str, Any]: + key = f"{run_prefix}/{payload_index:06d}.pt" + started = time.perf_counter() + store.put(key, body) + s3_put_s = time.perf_counter() - started + try: + started = time.perf_counter() + responses = post_vllm_refit_endpoints( + endpoint_urls, + { + "bucket": store.bucket, + "region": store.region, + "key": key, + "checksum": sparse_payload_checksum(body), + }, + api_key=api_key, + timeout_s=timeout_s, + ) + manifest_post_s = time.perf_counter() - started + finally: + with suppress(Exception): + store.delete(key) + return { + "s3_put_s": s3_put_s, + "manifest_post_s": manifest_post_s, + "receiver": merge_vllm_refit_receiver_timing({}, responses, maximum=True), + } + + return stream_sparse_delta_payloads( + iterator, + delta_tracker=delta_tracker, + transport="s3", + send_payload=send_payload, + transfer_workers=refit_env_int( + "NRL_REFIT_S3_UPLOAD_WORKERS", + default=max(4, min(32, os.cpu_count() or 32)), + ), + shard_rank=shard_rank, + shard_count=shard_count, + ) + + +def post_vllm_refit_endpoints( + endpoint_urls: Sequence[str], + body: Mapping[str, str] | bytes, + *, + api_key: str | None, + timeout_s: float, + headers: Mapping[str, str] | None = None, + executor: ThreadPoolExecutor | None = None, +) -> list[dict[str, Any]]: + request_headers = dict(headers or {}) + if api_key: + request_headers[G_VLLM_REFIT_API_KEY_HEADER] = api_key + request_kwargs: dict[str, Any] = ( + {"data": body} if isinstance(body, bytes) else {"json": body} + ) + + def post(url: str) -> dict[str, Any]: + response = refit_http_session().post( + url, + **request_kwargs, + headers=request_headers, + timeout=timeout_s, + ) + result: dict[str, Any] = response.json() if response.content else {} + if response.status_code >= 400 or result.get("ok") is not True: + raise RuntimeError(f"vLLM refit failed for {url}: {result}") + return result + + pool = executor or _executor("refit-fanout", len(endpoint_urls)) + futures = [pool.submit(post, url) for url in endpoint_urls] + wait(futures) + return [future.result() for future in futures] + + +def flush_vllm_refit_urls( + base_urls: Sequence[str], + *, + api_key_env_var: str | None, + timeout_s: float, +) -> list[dict[str, Any]]: + endpoint_urls = [ + f"{url}{G_VLLM_REFIT_FLUSH_PATH}" + for url in (url.strip().rstrip("/") for url in base_urls if url.strip()) + ] + return post_vllm_refit_endpoints( + endpoint_urls, + {}, + api_key=vllm_refit_api_key(api_key_env_var), + timeout_s=timeout_s, + ) + + +def download_s3_refit_payload( + manifest: Mapping[str, Any], +) -> bytes: + bucket, region, key, checksum = ( + str(manifest[field]) for field in ("bucket", "region", "key", "checksum") + ) + body = _get_manifest_s3_store(bucket, region).get(key) + return decode_sparse_payload(body, checksum) + + +def merge_vllm_refit_receiver_timing( + result: dict[str, Any], + timings: Iterable[Mapping[str, Any]], + *, + maximum: bool, +) -> dict[str, Any]: + for timing in timings: + for key, value in timing.items(): + if key.startswith("receiver_") and key.endswith("_s"): + number = float(value) + if key in result: + number = ( + max(float(result[key]), number) + if maximum + else float(result[key]) + number + ) + result[key] = number + return result + + +def zstd_compress(raw: bytes, threads_env: str) -> bytes: + compressor = getattr(_CONTROL_SESSION_LOCAL, "zstd_compressor", None) + if compressor is None: + compressor = zstandard.ZstdCompressor( + level=1, + threads=refit_env_int(threads_env, default=0, min_value=0), + ) + _CONTROL_SESSION_LOCAL.zstd_compressor = compressor + return compressor.compress(raw) diff --git a/nemo_rl/utils/weight_transfer_s3.py b/nemo_rl/utils/weight_transfer_s3.py deleted file mode 100644 index 1784e48ab0e..00000000000 --- a/nemo_rl/utils/weight_transfer_s3.py +++ /dev/null @@ -1,122 +0,0 @@ -# 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. - -"""AWS CRT object transport for S3 refit payloads.""" - -import io -from functools import cache -from typing import Any -from urllib.parse import quote - -_PART_SIZE = 64 * 1024**2 -_MEMORY_LIMIT = 2 * 1024**3 - - -@cache -def _s3_client(region: str) -> Any: - from awscrt.auth import AwsCredentialsProvider - from awscrt.io import ClientBootstrap, DefaultHostResolver, EventLoopGroup - from awscrt.s3 import S3Client, create_default_s3_signing_config - - event_loop_group = EventLoopGroup() - bootstrap = ClientBootstrap( - event_loop_group, - DefaultHostResolver(event_loop_group), - ) - credentials = AwsCredentialsProvider.new_default_chain(bootstrap) - return S3Client( - bootstrap=bootstrap, - region=region, - signing_config=create_default_s3_signing_config( - region=region, - credential_provider=credentials, - ), - part_size=_PART_SIZE, - multipart_upload_threshold=_PART_SIZE, - throughput_target_gbps=10.0, - memory_limit=_MEMORY_LIMIT, - ) - - -class S3ObjectStore: - """Blocking object operations backed by CRT's asynchronous S3 client.""" - - def __init__(self, *, bucket: str, region: str) -> None: - from awscrt.s3 import S3RequestType - - self.bucket = bucket - self.region = region - self._client = _s3_client(region) - self._request_type = S3RequestType - - def put_object(self, key: str, body: bytes) -> None: - request = self._request("PUT", key, body) - self._client.make_request( - type=self._request_type.PUT_OBJECT, - request=request, - ).finished_future.result() - - def get_object(self, key: str) -> bytes: - from awscrt.http import HttpHeaders - - body = bytearray() - - def on_headers( - status_code: int, - headers: list[tuple[str, str]], - **_kwargs: Any, - ) -> None: - nonlocal body - if status_code != 200: - raise RuntimeError(f"S3 GET returned HTTP {status_code}.") - length = HttpHeaders(headers).get("content-length") - if length is None: - raise RuntimeError("S3 GET response omitted content-length.") - body = bytearray(int(length)) - - def on_body(chunk: bytes, offset: int, **_kwargs: Any) -> None: - body[offset : offset + len(chunk)] = chunk - - self._client.make_request( - type=self._request_type.GET_OBJECT, - request=self._request("GET", key), - on_headers=on_headers, - on_body=on_body, - ).finished_future.result() - return bytes(body) - - def delete_object(self, key: str) -> None: - self._client.make_request( - type=self._request_type.DEFAULT, - request=self._request("DELETE", key), - operation_name="DeleteObject", - ).finished_future.result() - - def _request(self, method: str, key: str, body: bytes | None = None) -> Any: - from awscrt.http import HttpHeaders, HttpRequest - - headers = HttpHeaders( - [("host", f"{self.bucket}.s3.{self.region}.amazonaws.com")] - ) - if body is not None: - headers.add("content-length", str(len(body))) - headers.add("content-type", "application/octet-stream") - elif method == "DELETE": - headers.add("content-length", "0") - return HttpRequest( - method, - f"/{quote(key, safe='/~')}", - headers, - io.BytesIO(body) if body is not None else None, - ) diff --git a/nemo_rl/utils/weight_transfer_s3_manifest.py b/nemo_rl/utils/weight_transfer_s3_manifest.py deleted file mode 100644 index 3ad35aa058f..00000000000 --- a/nemo_rl/utils/weight_transfer_s3_manifest.py +++ /dev/null @@ -1,430 +0,0 @@ -# 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. - -"""S3 manifest control-plane helpers for sparse vLLM refit.""" - -import io -import os -import threading -import time -import uuid -from collections.abc import Iterable, Iterator, Mapping, Sequence -from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait -from contextlib import suppress -from functools import cache -from typing import Any - -import requests -import torch -import zstandard -from urllib3.util.retry import Retry - -from nemo_rl.utils.packed_tensor import get_target_packed_tensor_size -from nemo_rl.utils.weight_transfer_s3 import S3ObjectStore -from nemo_rl.utils.weight_transfer_sparse_codec import ( - DeltaCompressionTracker, - NamedTensor, - TensorBatch, -) - -G_VLLM_REFIT_S3_MANIFEST_PATH = "/nemo-rl/refit/s3-manifest" -G_VLLM_REFIT_FLUSH_PATH = "/nemo-rl/refit/flush" -G_VLLM_REFIT_API_KEY_HEADER = "x-nemo-rl-refit-key" -_CONTROL_SESSION_LOCAL = threading.local() - - -def _env_int(name: str, default: int, min_value: int = 1) -> int: - value = int(os.getenv(name) or default) - if value < min_value: - raise ValueError(f"{name} must be >= {min_value}.") - return value - - -def _iter_chunks( - tensors: Iterable[NamedTensor], target_bytes: int -) -> Iterator[tuple[TensorBatch, float]]: - iterator = iter(tensors) - pending = None - while True: - started = time.perf_counter() - chunk = [pending] if pending is not None else [] - size = pending[1].numel() * pending[1].element_size() if pending else 0 - pending = None - for item in iterator: - item_size = item[1].numel() * item[1].element_size() - if chunk and size + item_size > target_bytes: - pending = item - break - chunk.append(item) - size += item_size - if size >= target_bytes: - break - export_pull_s = time.perf_counter() - started - if not chunk: - return - yield chunk, export_pull_s - - -def _http_session() -> requests.Session: - session = getattr(_CONTROL_SESSION_LOCAL, "session", None) - if session is None: - session = requests.Session() - adapter = requests.adapters.HTTPAdapter( - pool_connections=64, - pool_maxsize=64, - max_retries=Retry( - total=3, - backoff_factor=0.25, - status_forcelist=(500, 502, 503, 504), - allowed_methods={"POST"}, - ), - ) - session.mount("http://", adapter) - session.mount("https://", adapter) - _CONTROL_SESSION_LOCAL.session = session - return session - - -@cache -def _get_manifest_s3_store(bucket: str, region: str) -> S3ObjectStore: - return S3ObjectStore(bucket=bucket, region=region) - - -def vllm_refit_api_key(api_key_env_var: str | None) -> str | None: - if not api_key_env_var: - return None - token = os.environ.get(api_key_env_var) - if not token: - raise RuntimeError( - "vLLM S3 refit API key env var " - f"{api_key_env_var!r} is configured but unset or empty." - ) - return token - - -def _s3_sparse_export_chunk_size(delta_tracker: DeltaCompressionTracker) -> int: - requested = _env_int( - "NRL_REFIT_S3_EXPORT_CHUNK_BYTES", - default=256 * 1024**2, - min_value=1, - ) - if torch.cuda.is_available(): - requested = min(requested, get_target_packed_tensor_size()) - return min(requested, delta_tracker.sparse_bucket_size_bytes) - - -@cache -def _executor(key: str, workers: int) -> ThreadPoolExecutor: - return ThreadPoolExecutor(max_workers=workers, thread_name_prefix=f"nrl-{key}") - - -def _require_delta_tracker( - delta_tracker: DeltaCompressionTracker | None, -) -> DeltaCompressionTracker: - if delta_tracker is None: - raise RuntimeError("vLLM S3 sparse refit requires delta compression.") - return delta_tracker - - -def init_sparse_delta_baseline_from_iterator( - iterator: Iterable[NamedTensor], - *, - delta_tracker: DeltaCompressionTracker | None, - shard_rank: int = 0, - shard_count: int = 1, -) -> None: - start_s = time.perf_counter() - delta_tracker = _require_delta_tracker(delta_tracker) - export_chunk_size = _s3_sparse_export_chunk_size(delta_tracker) - - chunk_count = 0 - export_pull_s = snapshot_s = 0.0 - for chunk_index, (chunk, pull_s) in enumerate( - _iter_chunks(iterator, export_chunk_size) - ): - chunk_count = chunk_index + 1 - export_pull_s += pull_s - if chunk_index % shard_count != shard_rank: - continue - started = time.perf_counter() - delta_tracker.snapshot_baseline(chunk) - snapshot_s += time.perf_counter() - started - print( - "REFIT_BASELINE_INIT " - f"event=end chunks={chunk_count} export_pull_s={export_pull_s:.3f} " - f"snapshot_s={snapshot_s:.3f} " - f"seconds={time.perf_counter() - start_s:.3f}", - flush=True, - ) - - -def stream_sparse_delta_payloads_via_s3_manifest( - iterator: Iterable[NamedTensor], - *, - delta_tracker: DeltaCompressionTracker | None, - refit_urls: Sequence[str], - api_key_env_var: str | None = None, - timeout_s: float = 600.0, - shard_rank: int = 0, - shard_count: int = 1, -) -> dict[str, Any]: - urls = [url.strip().rstrip("/") for url in refit_urls if url.strip()] - if not urls: - raise ValueError("At least one vLLM S3 refit URL is required.") - delta_tracker = _require_delta_tracker(delta_tracker) - - bucket = os.getenv("NRL_REFIT_S3_BUCKET", "").strip() - if not bucket: - raise RuntimeError("NRL_REFIT_S3_BUCKET must be set for S3 refit.") - store = _get_manifest_s3_store( - bucket, - os.getenv("NRL_REFIT_S3_REGION", "us-east-1").strip() or "us-east-1", - ) - endpoint_urls = [f"{url}{G_VLLM_REFIT_S3_MANIFEST_PATH}" for url in urls] - object_prefix = os.getenv("NRL_REFIT_S3_PREFIX", "nemo-rl-refit").strip("/") - run_prefix = ( - f"{object_prefix}/{uuid.uuid4().hex}" if object_prefix else uuid.uuid4().hex - ) - encode_workers = _env_int( - "NRL_REFIT_S3_ENCODE_WORKERS", - default=max(2, min(8, os.cpu_count() or 8)), - min_value=1, - ) - upload_workers = _env_int( - "NRL_REFIT_S3_UPLOAD_WORKERS", - default=max(4, min(32, os.cpu_count() or 32)), - min_value=1, - ) - pipeline_workers = max(encode_workers, upload_workers) - executor = _executor("refit-s3-pipeline", pipeline_workers) - encode_slots = threading.Semaphore(encode_workers) - export_chunk_size = _s3_sparse_export_chunk_size(delta_tracker) - - def process_chunk(chunk: TensorBatch, payload_index: int) -> dict[str, Any] | None: - with encode_slots: - started = time.perf_counter() - payload = delta_tracker.prepare_sparse_delta_payload(chunk) - encode_s = time.perf_counter() - started - if not payload[2]: - return None - started = time.perf_counter() - buffer = io.BytesIO() - torch.save(payload, buffer) - raw_body = buffer.getvalue() - serialize_s = time.perf_counter() - started - started = time.perf_counter() - body = _zstd_compress(raw_body) - compress_s = time.perf_counter() - started - - key = f"{run_prefix}/{payload_index:06d}.pt" - started = time.perf_counter() - store.put_object(key, body) - s3_put_s = time.perf_counter() - started - - manifest = { - "bucket": store.bucket, - "region": store.region, - "key": key, - } - try: - started = time.perf_counter() - responses = _post_refit_body_to_endpoint_urls( - endpoint_urls, - manifest, - api_key_env_var=api_key_env_var, - timeout_s=timeout_s, - ) - manifest_post_s = time.perf_counter() - started - finally: - with suppress(Exception): - store.delete_object(key) - - return { - "body_size": len(body), - "encode_s": encode_s, - "serialize_s": serialize_s, - "compress_s": compress_s, - "s3_put_s": s3_put_s, - "manifest_post_s": manifest_post_s, - "receiver": merge_vllm_refit_receiver_timing({}, responses, maximum=True), - } - - timing: dict[str, float] = {} - receiver_timing: dict[str, float] = {} - counts = {"payloads": 0, "uploaded_bytes": 0} - chunk_count = 0 - export_pull_s = 0.0 - inflight: set[Any] = set() - max_inflight = pipeline_workers * 2 - - def collect_completed(future: Any) -> None: - result = future.result() - if result is None: - return - counts["payloads"] += 1 - counts["uploaded_bytes"] += int(result["body_size"]) - for key, value in result.items(): - if key.endswith("_s"): - timing[key] = timing.get(key, 0.0) + float(value) - merge_vllm_refit_receiver_timing( - receiver_timing, [result["receiver"]], maximum=False - ) - - def drain_completed() -> None: - completed, _ = wait(inflight, return_when=FIRST_COMPLETED) - for future in completed: - inflight.remove(future) - collect_completed(future) - - payload_index = 0 - stream_start = time.perf_counter() - try: - for chunk_index, (chunk, pull_s) in enumerate( - _iter_chunks(iterator, export_chunk_size) - ): - chunk_count = chunk_index + 1 - export_pull_s += pull_s - if chunk_index % shard_count != shard_rank: - continue - if len(inflight) >= max_inflight: - drain_completed() - inflight.add(executor.submit(process_chunk, chunk, payload_index)) - payload_index += 1 - - while inflight: - drain_completed() - except Exception: - for future in inflight: - future.cancel() - wait(inflight) - with suppress(Exception): - flush_vllm_refit_urls( - urls, - api_key_env_var=api_key_env_var, - timeout_s=min(timeout_s, 60.0), - ) - raise - - timing = { - "total_s": time.perf_counter() - stream_start, - "export_pull_s": export_pull_s, - **timing, - "payloads": counts["payloads"], - "chunks": chunk_count, - "uploaded_mb": counts["uploaded_bytes"] / 1e6, - "pipeline_workers": pipeline_workers, - "encode_workers": encode_workers, - "export_chunk_mb": export_chunk_size / 1e6, - "shard_rank": shard_rank, - "shard_count": shard_count, - } - timing.update(receiver_timing) - print( - "REFIT_S3_TIMING " - + " ".join(f"{key}={value}" for key, value in timing.items()), - flush=True, - ) - return {"ok": True, "payloads": counts["payloads"]} - - -def _post_refit_body_to_endpoint_urls( - endpoint_urls: Sequence[str], - body: Mapping[str, str], - *, - api_key_env_var: str | None, - timeout_s: float, -) -> list[dict[str, Any]]: - headers = {} - if token := vllm_refit_api_key(api_key_env_var): - headers[G_VLLM_REFIT_API_KEY_HEADER] = token - - def post(url: str) -> dict[str, Any]: - response = _http_session().post( - url, - json=body, - headers=headers, - timeout=timeout_s, - ) - result: dict[str, Any] = response.json() if response.content else {} - if response.status_code >= 400 or result.get("ok") is not True: - raise RuntimeError(f"vLLM refit failed for {url}: {result}") - return result - - return list(_executor("refit-fanout", len(endpoint_urls)).map(post, endpoint_urls)) - - -def flush_vllm_refit_urls( - base_urls: Sequence[str], - *, - api_key_env_var: str | None, - timeout_s: float, -) -> None: - endpoint_urls = [ - f"{url}{G_VLLM_REFIT_FLUSH_PATH}" - for url in (url.strip().rstrip("/") for url in base_urls if url.strip()) - ] - _post_refit_body_to_endpoint_urls( - endpoint_urls, - {}, - api_key_env_var=api_key_env_var, - timeout_s=timeout_s, - ) - - -def download_s3_refit_payload( - manifest: Mapping[str, Any], -) -> bytes: - bucket, region, key = ( - str(manifest[field]) for field in ("bucket", "region", "key") - ) - return _zstd_decompress(_get_manifest_s3_store(bucket, region).get_object(key)) - - -def merge_vllm_refit_receiver_timing( - result: dict[str, Any], - timings: Iterable[Mapping[str, Any]], - *, - maximum: bool, -) -> dict[str, Any]: - for timing in timings: - for key, value in timing.items(): - if key.startswith("receiver_") and key.endswith("_s"): - number = float(value) - if key in result: - number = ( - max(float(result[key]), number) - if maximum - else float(result[key]) + number - ) - result[key] = number - return result - - -def _zstd_compress(raw: bytes) -> bytes: - compressor = getattr(_CONTROL_SESSION_LOCAL, "zstd_compressor", None) - if compressor is None: - compressor = zstandard.ZstdCompressor( - level=1, - threads=_env_int("NRL_REFIT_S3_ZSTD_THREADS", default=0, min_value=0), - ) - _CONTROL_SESSION_LOCAL.zstd_compressor = compressor - return compressor.compress(raw) - - -def _zstd_decompress(raw: bytes) -> bytes: - decompressor = getattr(_CONTROL_SESSION_LOCAL, "zstd_decompressor", None) - if decompressor is None: - decompressor = zstandard.ZstdDecompressor() - _CONTROL_SESSION_LOCAL.zstd_decompressor = decompressor - return decompressor.decompress(raw) diff --git a/nemo_rl/utils/weight_transfer_sparse_codec.py b/nemo_rl/utils/weight_transfer_sparse_codec.py index a0cafc5d8c4..2786f1ef527 100644 --- a/nemo_rl/utils/weight_transfer_sparse_codec.py +++ b/nemo_rl/utils/weight_transfer_sparse_codec.py @@ -25,6 +25,7 @@ NamedTensor = tuple[str, torch.Tensor] TensorBatch = list[NamedTensor] TensorPayload = tuple[torch.Tensor, torch.Tensor, list[dict[str, Any]]] +PreparedTensorPayload = tuple[TensorPayload, int, int] def encode_sparse_infos( @@ -122,63 +123,106 @@ def __init__(self, config: Mapping[str, Any]) -> None: self.sparse_bucket_size_bytes = int(config["sparse_bucket_size_bytes"]) if self.sparse_bucket_size_bytes < 1: raise ValueError("delta_compression.sparse_bucket_size_bytes must be >= 1") - dtype_name = {"bf16": "bfloat16", "fp16": "float16", "fp32": "float32"}.get( - str(config["dtype"]).lower(), str(config["dtype"]).lower() + self.delta_dtype = { + "bf16": torch.bfloat16, + "bfloat16": torch.bfloat16, + "fp16": torch.float16, + "float16": torch.float16, + "fp32": torch.float32, + "float32": torch.float32, + }[str(config["dtype"]).lower()] + self.verification_samples = int( + os.getenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "0") ) - self.delta_dtype: torch.dtype = getattr(torch, dtype_name) + if self.verification_samples < 0: + raise ValueError("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD must be >= 0") self.baseline_in_memory = os.getenv("NRL_REFIT_BASELINE_IN_MEMORY") == "1" - self.baseline_mmap_dir = config.get("baseline_mmap_dir") or os.getenv( - "NRL_REFIT_BASELINE_MMAP_DIR" - ) + self.baseline_mmap_dir = os.getenv("NRL_REFIT_BASELINE_MMAP_DIR") self.baseline: dict[str, torch.Tensor] = {} self._pending_updates: dict[str, tuple[torch.Tensor, torch.Tensor]] = {} self._pending_updates_lock = threading.Lock() self._baseline_commits: tuple[Any, ...] = () - self._baseline_commit_lock = threading.Lock() self._baseline_commit_executor = ThreadPoolExecutor( max_workers=4, thread_name_prefix="nrl-refit-baseline" ) - def prepare_sparse_delta_payload(self, tensors: TensorBatch) -> TensorPayload: + def prepare_sparse_delta_payload( + self, tensors: TensorBatch + ) -> PreparedTensorPayload: self._wait_for_baseline_commits() sparse_infos = [] + verification_sources = [] pending_updates = {} + changed_elements = total_elements = 0 for name, tensor in tensors: baseline = self.baseline.get(name) if baseline is None: raise RuntimeError(f"Sparse delta baseline is missing {name!r}.") current = tensor.detach().cpu() current_flat, baseline_flat = current.view(-1), baseline.view(-1) + total_elements += current_flat.numel() locations = (current_flat != baseline_flat).nonzero().view(-1) + changed_elements += locations.numel() if locations.numel(): current_values = current_flat[locations] + baseline_values = baseline_flat[locations] + deltas = (current_values - baseline_values).to(self.delta_dtype) + expected_values = baseline_values + deltas.to(baseline.dtype) sparse_infos.append( ( name, current, locations, - (current_values - baseline_flat[locations]).to( - self.delta_dtype - ), + deltas, ) ) - pending_updates[name] = (locations, current_values) + if self.verification_samples: + verification_sources.append((locations, deltas)) + pending_updates[name] = (locations, expected_values) with self._pending_updates_lock: self._pending_updates.update(pending_updates) - return encode_sparse_infos(sparse_infos, empty_dtype=self.delta_dtype) + payload = encode_sparse_infos(sparse_infos, empty_dtype=self.delta_dtype) + if verification_sources: + self._add_verification_samples(payload[2], verification_sources) + return payload, changed_elements, total_elements + + def _add_verification_samples( + self, + metadata: list[dict[str, Any]], + sources: list[tuple[torch.Tensor, torch.Tensor]], + ) -> None: + total = sum(int(locations.numel()) for locations, _ in sources) + count = min(self.verification_samples, total) + if not count: + return + + sample_ranks = [ + ((2 * index + 1) * total) // (2 * count) for index in range(count) + ] + sample_index = offset = 0 + for item, (locations, deltas) in zip(metadata, sources, strict=True): + end = offset + locations.numel() + while sample_index < count and sample_ranks[sample_index] < end: + local_index = sample_ranks[sample_index] - offset + location = int(locations[local_index]) + item.setdefault("verification_locations", []).append(location) + item.setdefault("verification_deltas", []).append( + float(deltas[local_index]) + ) + sample_index += 1 + offset = end def on_sync_succeeded(self) -> None: with self._pending_updates_lock: pending_updates, self._pending_updates = self._pending_updates, {} items = list(pending_updates.items()) workers = min(4, len(items)) - with self._baseline_commit_lock: - self._baseline_commits = tuple( - self._baseline_commit_executor.submit( - self._commit_baseline_updates, items[worker::workers] - ) - for worker in range(workers) + self._baseline_commits = tuple( + self._baseline_commit_executor.submit( + self._commit_baseline_updates, items[worker::workers] ) + for worker in range(workers) + ) def on_sync_failed(self) -> None: with self._pending_updates_lock: @@ -190,13 +234,9 @@ def snapshot_baseline(self, tensors: Iterable[NamedTensor]) -> None: self._baseline(name, tuple(tensor.shape), tensor.dtype).copy_(tensor) def _wait_for_baseline_commits(self) -> None: - with self._baseline_commit_lock: - commits = self._baseline_commits - for commit in commits: + for commit in self._baseline_commits: commit.result() - with self._baseline_commit_lock: - if self._baseline_commits == commits: - self._baseline_commits = () + self._baseline_commits = () def _commit_baseline_updates( self, updates: Iterable[tuple[str, tuple[torch.Tensor, torch.Tensor]]] diff --git a/nemo_rl/utils/weight_transfer_zmq.py b/nemo_rl/utils/weight_transfer_zmq.py new file mode 100644 index 00000000000..f43999a9f83 --- /dev/null +++ b/nemo_rl/utils/weight_transfer_zmq.py @@ -0,0 +1,406 @@ +# 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. + +"""Transactional ZeroMQ value plane for remote sparse vLLM refit.""" + +import json +import threading +import time +import uuid +from collections.abc import Iterable, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from typing import Any + +import zmq + +from nemo_rl.utils.weight_transfer_remote_sparse import ( + SparseDeltaStreamResult, + merge_vllm_refit_receiver_timing, + post_vllm_refit_endpoints, + refit_env_int, + sparse_payload_checksum, + stream_sparse_delta_payloads, + vllm_refit_api_key, +) +from nemo_rl.utils.weight_transfer_sparse_codec import ( + DeltaCompressionTracker, + NamedTensor, +) + +G_VLLM_REFIT_ZMQ_PAYLOAD_PATH = "/nemo-rl/refit/zmq-payload" +G_VLLM_REFIT_TRANSFER_HEADER = "x-nemo-rl-refit-transfer" +G_VLLM_REFIT_PRODUCER_HEADER = "x-nemo-rl-refit-producer" +G_VLLM_REFIT_PAYLOAD_HEADER = "x-nemo-rl-refit-payload" +G_VLLM_REFIT_CHECKSUM_HEADER = "x-nemo-rl-refit-checksum" + +_PROTOCOL = "nemo-rl-sparse-zmq-v1" +_DATA = b"DATA" +_ACK = b"ACK" +_NACK = b"NACK" +_ZMQ_LOCAL = threading.local() + + +def _json_bytes(value: Mapping[str, Any]) -> bytes: + return json.dumps(value, separators=(",", ":"), sort_keys=True).encode() + + +class ZmqSparseRefitClient: + """One-thread DEALER client with retry-safe payload identifiers.""" + + def __init__( + self, + address: str, + *, + timeout_s: float, + producer_id: int, + api_key: str | None = None, + ) -> None: + self._address = address + self._timeout_ms = max(1, int(timeout_s * 1000)) + self._producer_id = producer_id + self._api_key = api_key + self._socket = zmq.Context.instance().socket(zmq.DEALER) + self._socket.setsockopt( + zmq.IDENTITY, + f"nrl-{producer_id}-{uuid.uuid4().hex}".encode(), + ) + self._socket.setsockopt(zmq.LINGER, 0) + self._socket.setsockopt(zmq.IMMEDIATE, 1) + self._socket.setsockopt(zmq.SNDHWM, 2) + self._socket.setsockopt(zmq.RCVHWM, 2) + self._socket.setsockopt(zmq.SNDTIMEO, self._timeout_ms) + self._socket.setsockopt(zmq.TCP_KEEPALIVE, 1) + self._socket.connect(address) + + def send_payload( + self, + *, + transfer_id: str, + payload_id: int, + checksum: str, + body: bytes, + ) -> dict[str, Any]: + metadata = { + "protocol": _PROTOCOL, + "transfer_id": transfer_id, + "producer_id": self._producer_id, + "payload_id": payload_id, + "checksum": checksum, + } + if self._api_key is not None: + metadata["api_key"] = self._api_key + metadata_frame = _json_bytes(metadata) + retries = refit_env_int("NRL_REFIT_ZMQ_RETRIES", default=3, min_value=0) + + for attempt in range(retries + 1): + try: + self._socket.send_multipart( + [_DATA, metadata_frame, body], + copy=False, + ) + except zmq.Again: + if attempt == retries: + break + continue + + deadline = time.monotonic() + self._timeout_ms / 1000 + while True: + remaining_ms = max(1, int((deadline - time.monotonic()) * 1000)) + if deadline <= time.monotonic() or not self._socket.poll( + remaining_ms, zmq.POLLIN + ): + break + frames = self._socket.recv_multipart() + if len(frames) != 2: + continue + kind, raw_reply = frames + reply = json.loads(raw_reply) + if kind == _NACK and "transfer_id" not in reply: + raise RuntimeError(f"ZeroMQ sparse refit rejected payload: {reply}") + reply_key = ( + reply.get("transfer_id"), + reply.get("producer_id"), + reply.get("payload_id"), + ) + if reply_key != (transfer_id, self._producer_id, payload_id): + continue + if kind == _ACK and reply.get("ok") is True: + return reply + raise RuntimeError(f"ZeroMQ sparse refit rejected payload: {reply}") + + if attempt < retries: + time.sleep(min(0.05 * 2**attempt, 0.5)) + + raise TimeoutError( + f"Timed out sending sparse refit payload {payload_id} to {self._address}." + ) + + def close(self) -> None: + self._socket.close() + + +class ZmqSparseRefitServer: + """Bounded ROUTER relay that fans each compressed payload to all replicas.""" + + def __init__( + self, + refit_urls: Sequence[str], + *, + bind_address: str, + api_key_env_var: str | None, + timeout_s: float, + ) -> None: + self._refit_endpoints = tuple( + f"{url}{G_VLLM_REFIT_ZMQ_PAYLOAD_PATH}" + for url in dict.fromkeys( + url.strip().rstrip("/") for url in refit_urls if url.strip() + ) + ) + if not self._refit_endpoints: + raise ValueError("ZeroMQ sparse refit requires receiver HTTP URLs.") + self._bind_address = bind_address + self._token = vllm_refit_api_key(api_key_env_var) + self._timeout_s = timeout_s + self._stop = threading.Event() + self._ready = threading.Event() + self._thread: threading.Thread | None = None + self._endpoint: str | None = None + self._error: Exception | None = None + self._payload_workers = refit_env_int( + "NRL_REFIT_ZMQ_RELAY_PAYLOAD_WORKERS", default=16 + ) + self._fanout_workers = refit_env_int( + "NRL_REFIT_ZMQ_RELAY_FANOUT_WORKERS", + default=max(8, min(32, len(self._refit_endpoints) * self._payload_workers)), + ) + + def start(self) -> str: + self._thread = threading.Thread( + target=self._run, + name="nrl-zmq-refit-relay", + daemon=True, + ) + self._thread.start() + if not self._ready.wait(timeout=10.0): + raise RuntimeError("Timed out starting the ZeroMQ sparse refit relay.") + if self._error is not None: + raise RuntimeError( + "Failed to start the ZeroMQ sparse refit relay." + ) from self._error + assert self._endpoint is not None + return self._endpoint + + def close(self) -> None: + self._stop.set() + if self._thread is not None: + self._thread.join(timeout=max(5.0, self._timeout_s)) + if self._thread.is_alive(): + raise RuntimeError("Timed out stopping the ZeroMQ sparse refit relay.") + self._thread = None + + def _fanout( + self, + body: bytes, + metadata: Mapping[str, Any], + http_executor: ThreadPoolExecutor, + ) -> dict[str, Any]: + headers = { + "content-type": "application/octet-stream", + G_VLLM_REFIT_TRANSFER_HEADER: str(metadata["transfer_id"]), + G_VLLM_REFIT_PRODUCER_HEADER: str(metadata["producer_id"]), + G_VLLM_REFIT_PAYLOAD_HEADER: str(metadata["payload_id"]), + G_VLLM_REFIT_CHECKSUM_HEADER: str(metadata["checksum"]), + } + started = time.perf_counter() + results = post_vllm_refit_endpoints( + self._refit_endpoints, + body, + api_key=self._token, + timeout_s=self._timeout_s, + headers=headers, + executor=http_executor, + ) + merged = merge_vllm_refit_receiver_timing({}, results, maximum=True) + merged["receiver_relay_fanout_s"] = time.perf_counter() - started + return merged + + @staticmethod + def _send_reply( + socket: zmq.Socket, + identity: bytes, + kind: bytes, + reply: Mapping[str, Any], + ) -> None: + try: + socket.send_multipart( + [identity, kind, _json_bytes(reply)], flags=zmq.NOBLOCK + ) + except (zmq.Again, zmq.ZMQError): + pass + + def _parse_data_message( + self, + frames: list[bytes], + ) -> tuple[bytes, tuple[str, int, int], bytes, dict[str, Any]]: + if len(frames) != 4: + raise ValueError(f"Expected 4 ZeroMQ frames, received {len(frames)}.") + identity, kind, raw_metadata, body = frames + if kind != _DATA: + raise ValueError(f"Unsupported ZeroMQ sparse refit message {kind!r}.") + metadata = json.loads(raw_metadata) + if metadata.get("protocol") != _PROTOCOL: + raise ValueError("Unsupported ZeroMQ sparse refit protocol.") + if self._token is not None and metadata.get("api_key") != self._token: + raise PermissionError("ZeroMQ sparse refit producer authentication failed.") + transfer_id = str(metadata["transfer_id"]) + producer_id = int(metadata["producer_id"]) + payload_id = int(metadata["payload_id"]) + checksum = str(metadata["checksum"]) + if not transfer_id or producer_id < 0 or payload_id < 0: + raise ValueError("Invalid ZeroMQ sparse refit payload identity.") + actual = sparse_payload_checksum(body) + if actual != checksum: + raise ValueError( + f"Sparse refit payload checksum mismatch: expected={checksum}, actual={actual}." + ) + return identity, (transfer_id, producer_id, payload_id), body, metadata + + def _run(self) -> None: + context = zmq.Context() + socket = context.socket(zmq.ROUTER) + payload_executor = ThreadPoolExecutor( + max_workers=self._payload_workers, + thread_name_prefix="nrl-zmq-payload", + ) + http_executor = ThreadPoolExecutor( + max_workers=self._fanout_workers, + thread_name_prefix="nrl-zmq-fanout", + ) + pending: dict[Any, tuple[bytes, tuple[str, int, int]]] = {} + try: + socket.setsockopt(zmq.LINGER, 0) + socket.setsockopt(zmq.ROUTER_MANDATORY, 1) + socket.setsockopt(zmq.SNDHWM, 16) + socket.setsockopt(zmq.RCVHWM, 16) + socket.setsockopt(zmq.TCP_KEEPALIVE, 1) + socket.bind(self._bind_address) + self._endpoint = socket.getsockopt_string(zmq.LAST_ENDPOINT) + self._ready.set() + + while not self._stop.is_set() or pending: + if not self._stop.is_set() and len(pending) < self._payload_workers: + if socket.poll(10, zmq.POLLIN): + frames = socket.recv_multipart() + identity = frames[0] if frames else b"" + try: + identity, key, body, metadata = self._parse_data_message( + frames + ) + future = payload_executor.submit( + self._fanout, + body, + metadata, + http_executor, + ) + pending[future] = (identity, key) + except Exception as exc: + self._send_reply( + socket, + identity, + _NACK, + {"ok": False, "error": str(exc)}, + ) + + for future, (identity, key) in list(pending.items()): + if not future.done(): + continue + transfer_id, producer_id, payload_id = key + reply: dict[str, Any] = { + "ok": True, + "transfer_id": transfer_id, + "producer_id": producer_id, + "payload_id": payload_id, + } + try: + reply.update(future.result()) + except Exception as exc: + reply.update(ok=False, error=str(exc)) + kind = _NACK + else: + kind = _ACK + self._send_reply(socket, identity, kind, reply) + del pending[future] + except Exception as exc: + self._error = exc + self._ready.set() + finally: + payload_executor.shutdown(wait=True, cancel_futures=True) + http_executor.shutdown(wait=True, cancel_futures=True) + socket.close() + context.term() + + +def stream_sparse_delta_payloads_via_zmq( + iterator: Iterable[NamedTensor], + *, + delta_tracker: DeltaCompressionTracker, + refit_targets: Sequence[str], + transfer_id: str, + api_key_env_var: str | None, + timeout_s: float, + shard_rank: int, + shard_count: int, +) -> SparseDeltaStreamResult: + addresses = [address.strip() for address in refit_targets if address.strip()] + if not addresses: + raise ValueError("At least one ZeroMQ sparse refit address is required.") + address = addresses[shard_rank % len(addresses)] + api_key = vllm_refit_api_key(api_key_env_var) + + def send_payload(body: bytes, payload_id: int) -> dict[str, Any]: + clients = getattr(_ZMQ_LOCAL, "clients", None) + if clients is None: + clients = {} + _ZMQ_LOCAL.clients = clients + client_key = (address, shard_rank, api_key) + client = clients.get(client_key) + if client is None: + client = ZmqSparseRefitClient( + address, + timeout_s=timeout_s, + producer_id=shard_rank, + api_key=api_key, + ) + clients[client_key] = client + started = time.perf_counter() + reply = client.send_payload( + transfer_id=transfer_id, + payload_id=payload_id, + checksum=sparse_payload_checksum(body), + body=body, + ) + return { + "zmq_send_s": time.perf_counter() - started, + "receiver": reply, + } + + return stream_sparse_delta_payloads( + iterator, + delta_tracker=delta_tracker, + transport="zmq", + send_payload=send_payload, + transfer_workers=refit_env_int("NRL_REFIT_ZMQ_SEND_WORKERS", default=4), + shard_rank=shard_rank, + shard_count=shard_count, + ) diff --git a/nemo_rl/weight_sync/interfaces.py b/nemo_rl/weight_sync/interfaces.py index e60317be36d..f0e0817d63a 100644 --- a/nemo_rl/weight_sync/interfaces.py +++ b/nemo_rl/weight_sync/interfaces.py @@ -62,7 +62,7 @@ def sync_weights( *, timer: Optional[Timer] = None, kv_scales: Optional[dict[str, float]] = None, - ) -> None: + ) -> Optional[dict[str, float]]: """Transfer the latest policy weights to the generation backend. This method encapsulates the full sync lifecycle: @@ -88,6 +88,9 @@ def sync_weights( which forwards them to ``policy.broadcast_weights_for_collective()``. IPC and HTTP transports ignore this parameter. + Returns: + Optional transport-specific scalar metrics for the current sync. + Raises: RuntimeError: If the weight transfer fails. """ diff --git a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py new file mode 100644 index 00000000000..f89adfa0267 --- /dev/null +++ b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py @@ -0,0 +1,221 @@ +# 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. + +"""Shared S3/ZeroMQ sparse synchronizer for remote non-colocated vLLM refit.""" + +import time +import uuid +from contextlib import nullcontext, suppress +from typing import Any + +import ray + +from nemo_rl.utils.timer import Timer +from nemo_rl.utils.weight_transfer_remote_sparse import flush_vllm_refit_urls +from nemo_rl.weight_sync.interfaces import WeightSynchronizer + + +class VllmRemoteSparseWeightSynchronizer(WeightSynchronizer): + def __init__( + self, + policy: Any, + generation: Any, + *, + transport: str, + api_key_env_var: str | None = None, + request_timeout_s: float = 600.0, + ) -> None: + self._policy = policy + self._generation = generation + self._transport = transport + self._refit_urls: list[str] = [] + self._targets: list[str] = [] + self._api_key_env_var = api_key_env_var + self._request_timeout_s = request_timeout_s + self._stale = True + self._baseline_init_refs: list[Any] | None = None + self._baseline_commit_refs: list[Any] | None = None + + def sync_weights( + self, + *, + timer: Timer | None = None, + kv_scales: dict[str, float] | None = None, + ) -> dict[str, float]: + timer_context = ( + timer.time("prepare_for_generation/transfer_and_update_weights") + if timer is not None + else nullcontext() + ) + with timer_context: + if self._baseline_commit_refs is not None: + ray.get(self._baseline_commit_refs) + self._baseline_commit_refs = None + if not self._generation.invalidate_kv_cache(): + raise RuntimeError( + f"vLLM KV cache invalidation failed before {self._transport} " + "weight update." + ) + + if self._baseline_init_refs is not None: + ray.get(self._baseline_init_refs) + self._baseline_init_refs = None + succeeded = False + try: + transfer_id = uuid.uuid4().hex + refs = self._policy.stream_remote_sparse_weights( + self._transport, + self._targets, + transfer_id=transfer_id, + api_key_env_var=self._api_key_env_var, + timeout_s=self._request_timeout_s, + ) + results = ray.get(refs) + payloads = sum(result["payloads"] for result in results) + changed_elements = sum(result["changed_elements"] for result in results) + total_elements = sum(result["total_elements"] for result in results) + changed_pct = 100.0 * changed_elements / max(total_elements, 1) + print( + f"REFIT_{self._transport.upper()}_DELTA_CHANGE " + f"changed_elements={changed_elements} " + f"total_elements={total_elements} " + f"changed_pct={changed_pct:.8g}", + flush=True, + ) + candidates = 0 + samples = 0 + exact_mismatches = 0 + mismatches = 0 + abs_sum = 0.0 + max_abs = 0.0 + commit_s = 0.0 + if payloads: + started = time.perf_counter() + receiver_results = flush_vllm_refit_urls( + self._refit_urls, + api_key_env_var=self._api_key_env_var, + timeout_s=self._request_timeout_s, + ) + candidates = sum( + int(result.get("verification_candidates", 0)) + for result in receiver_results + ) + samples = sum( + int(result.get("verification_samples", 0)) + for result in receiver_results + ) + exact_mismatches = sum( + int(result.get("verification_exact_mismatches", 0)) + for result in receiver_results + ) + mismatches = sum( + int(result.get("verification_mismatches", 0)) + for result in receiver_results + ) + abs_sum = sum( + float(result.get("verification_abs_sum", 0.0)) + for result in receiver_results + ) + max_abs = max( + ( + float(result.get("verification_max_abs", 0.0)) + for result in receiver_results + ), + default=0.0, + ) + if candidates or samples: + print( + f"REFIT_{self._transport.upper()}_DELTA_VERIFY " + f"candidates={candidates} samples={samples} " + f"exact_mismatches={exact_mismatches} " + f"mismatches={mismatches} " + f"mean_abs={abs_sum / max(samples, 1):.8g} " + f"max_abs={max_abs:.8g}", + flush=True, + ) + if mismatches: + raise RuntimeError( + f"Sparse refit sampled {mismatches} mismatched deltas " + f"out of {samples}." + ) + commit_s = time.perf_counter() - started + print( + f"REFIT_{self._transport.upper()}_GLOBAL_COMMIT " + f"transfer_id={transfer_id} payloads={payloads} " + f"seconds={commit_s:.3f}", + flush=True, + ) + succeeded = True + finally: + if not succeeded: + with suppress(Exception): + flush_vllm_refit_urls( + self._refit_urls, + api_key_env_var=self._api_key_env_var, + timeout_s=min(self._request_timeout_s, 60.0), + ) + self._baseline_commit_refs = ( + self._policy.finish_remote_sparse_delta_sync(succeeded) + ) + self._stale = False + return { + "delta/changed_elements": float(changed_elements), + "delta/total_elements": float(total_elements), + "delta/changed_pct": changed_pct, + "delta_verify/candidates": float(candidates), + "delta_verify/samples": float(samples), + "delta_verify/exact_mismatches": float(exact_mismatches), + "delta_verify/mismatches": float(mismatches), + "delta_verify/mismatch_pct": 100.0 * mismatches / max(samples, 1), + "delta_verify/mean_abs": abs_sum / max(samples, 1), + "delta_verify/max_abs": max_abs, + "transfer/payloads": float(payloads), + "transfer/global_commit_s": commit_s, + } + + @property + def is_stale(self) -> bool: + return self._stale + + def mark_stale(self) -> None: + self._stale = True + + def init_communicator(self) -> None: + self._baseline_init_refs = self._policy.init_remote_sparse_delta_baseline( + self._transport + ) + self._refit_urls = self._generation.report_refit_server_base_urls() + self._targets = self._refit_urls + if self._transport == "zmq": + self._targets = self._generation.start_zmq_sparse_refit_relays( + self._refit_urls + ) + if not self._refit_urls or not self._targets: + raise ValueError( + f"vLLM {self._transport} sparse refit endpoints are missing." + ) + self._stale = False + + def shutdown(self) -> None: + for ref in (self._baseline_init_refs or []) + ( + self._baseline_commit_refs or [] + ): + ray.cancel(ref, force=False) + if self._transport == "zmq": + self._generation.stop_zmq_sparse_refit_relays() + self._baseline_init_refs = None + self._baseline_commit_refs = None + self._refit_urls = [] + self._targets = [] + self._stale = True diff --git a/nemo_rl/weight_sync/vllm_s3_sparse_weight_synchronizer.py b/nemo_rl/weight_sync/vllm_s3_sparse_weight_synchronizer.py deleted file mode 100644 index 3c3c36109db..00000000000 --- a/nemo_rl/weight_sync/vllm_s3_sparse_weight_synchronizer.py +++ /dev/null @@ -1,122 +0,0 @@ -# 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. - -"""S3 manifest weight synchronizer for remote non-colocated vLLM refit.""" - -import time -from contextlib import nullcontext -from typing import Any - -import ray - -from nemo_rl.utils.timer import Timer -from nemo_rl.utils.weight_transfer_s3_manifest import flush_vllm_refit_urls -from nemo_rl.weight_sync.interfaces import WeightSynchronizer - - -class VllmS3SparseWeightSynchronizer(WeightSynchronizer): - def __init__( - self, - policy: Any, - generation: Any, - *, - api_key_env_var: str | None = None, - request_timeout_s: float = 600.0, - ) -> None: - self._policy = policy - self._generation = generation - self._refit_urls: list[str] = [] - self._api_key_env_var = api_key_env_var - self._request_timeout_s = request_timeout_s - self._stale = True - self._baseline_init_refs: list[Any] | None = None - self._baseline_commit_refs: list[Any] | None = None - - def sync_weights( - self, - *, - timer: Timer | None = None, - kv_scales: dict[str, float] | None = None, - ) -> None: - timer_context = ( - timer.time("prepare_for_generation/transfer_and_update_weights") - if timer is not None - else nullcontext() - ) - with timer_context: - if self._baseline_commit_refs is not None: - ray.get(self._baseline_commit_refs) - self._baseline_commit_refs = None - flush_success = self._generation.invalidate_kv_cache() - if not flush_success: - print("vLLM KV cache invalidation failed before S3 weight update.") - - if self._baseline_init_refs is not None: - ray.get(self._baseline_init_refs) - self._baseline_init_refs = None - succeeded = False - try: - results = ray.get( - self._policy.stream_sparse_weights_via_s3_manifest( - self._refit_urls, - api_key_env_var=self._api_key_env_var, - timeout_s=self._request_timeout_s, - ) - ) - payloads = sum(int(result["payloads"]) for result in results) - if payloads: - started = time.perf_counter() - flush_vllm_refit_urls( - self._refit_urls, - api_key_env_var=self._api_key_env_var, - timeout_s=self._request_timeout_s, - ) - print( - "REFIT_S3_GLOBAL_FLUSH " - f"payloads={payloads} " - f"seconds={time.perf_counter() - started:.3f}", - flush=True, - ) - succeeded = True - finally: - self._baseline_commit_refs = ( - self._policy.finish_remote_sparse_delta_sync(succeeded) - ) - self._stale = False - - @property - def is_stale(self) -> bool: - return self._stale - - def mark_stale(self) -> None: - self._stale = True - - def init_communicator(self) -> None: - self._baseline_init_refs = self._policy.init_remote_sparse_delta_baseline() - self._refit_urls = self._generation.report_refit_server_base_urls() - if not self._refit_urls: - raise ValueError( - "vLLM S3 sparse refit requires expose_http_refit_server=true." - ) - self._stale = False - - def shutdown(self) -> None: - for ref in (self._baseline_init_refs or []) + ( - self._baseline_commit_refs or [] - ): - ray.cancel(ref, force=False) - self._baseline_init_refs = None - self._baseline_commit_refs = None - self._refit_urls = [] - self._stale = True diff --git a/pyrefly.toml b/pyrefly.toml index 2a304c21672..3aece31741a 100644 --- a/pyrefly.toml +++ b/pyrefly.toml @@ -200,10 +200,10 @@ project-includes = [ "nemo_rl/weight_sync/http_weight_synchronizer.py", "nemo_rl/weight_sync/interfaces.py", "nemo_rl/weight_sync/ipc_weight_synchronizer.py", - "nemo_rl/weight_sync/vllm_s3_sparse_weight_synchronizer.py", - "nemo_rl/utils/weight_transfer_s3_manifest.py", - "nemo_rl/utils/weight_transfer_s3.py", + "nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py", + "nemo_rl/utils/weight_transfer_remote_sparse.py", "nemo_rl/utils/weight_transfer_sparse_codec.py", + "nemo_rl/utils/weight_transfer_zmq.py", "tools/model_diagnostics/1.max_model_len_respected.py", "tools/model_diagnostics/2.long_generation_decode_vs_prefill.py", "tools/model_diagnostics/3.check_and_reinit_hf_model_embeddings_untrained.py", diff --git a/tests/unit/algorithms/test_grpo.py b/tests/unit/algorithms/test_grpo.py index df91ba56a0f..bc8833f1fd0 100644 --- a/tests/unit/algorithms/test_grpo.py +++ b/tests/unit/algorithms/test_grpo.py @@ -33,6 +33,7 @@ _apply_mask_sample_filter, _apply_message_level_advantage_penalties, _default_grpo_save_state, + _initial_policy_generation_stale, _raise_if_reward_penalties_enabled_without_nemo_gym, _resolve_message_level_advantage_penalties, _should_use_async_rollouts, @@ -115,6 +116,17 @@ def test_missing_mask_sample_is_noop(self): ) +def test_initial_policy_generation_stale() -> None: + generation = MagicMock() + generation.weight_synchronizer.is_stale = False + + assert not _initial_policy_generation_stale(generation, completed_steps=0) + assert _initial_policy_generation_stale(generation, completed_steps=1) + + generation.weight_synchronizer.is_stale = True + assert _initial_policy_generation_stale(generation, completed_steps=0) + + @pytest.fixture def mock_grpo_components(): # Create mock components @@ -1863,7 +1875,9 @@ def fake_batched_message_log_to_flat_message(*_args, **_kwargs): lambda *_args, **_kwargs: (torch.tensor([0.1]), torch.tensor([1.0])), ) monkeypatch.setattr( - grpo_mod, "refit_policy_generation", lambda *_args, **_kwargs: None + grpo_mod, + "refit_policy_generation", + lambda *_args, **_kwargs: {"delta/changed_pct": 4.0}, ) monkeypatch.setattr( grpo_mod, "print_performance_metrics", lambda *_args, **_kwargs: {} @@ -1877,12 +1891,25 @@ def fake_batched_message_log_to_flat_message(*_args, **_kwargs): lambda *_args, **_kwargs: seq_logprob_error_result, ) + dynamic_sampling_calls = 0 + + def fake_dynamic_sampling(repeated_batch, *_args, **_kwargs): + nonlocal dynamic_sampling_calls + dynamic_sampling_calls += 1 + repeated_batch["filtered_reward"] = repeated_batch["total_reward"] + repeated_batch["baseline"] = torch.zeros(repeated_batch.size) + repeated_batch["std"] = torch.ones(repeated_batch.size) + complete = dynamic_sampling_calls == 2 + return repeated_batch, complete, None if complete else repeated_batch, {} + + monkeypatch.setattr(grpo_mod, "dynamic_sampling", fake_dynamic_sampling) + master_config = mock_grpo_components["master_config"] master_config.grpo["max_num_steps"] = 1 master_config.grpo["max_num_epochs"] = 1 master_config.grpo["val_period"] = 0 master_config.grpo["val_at_start"] = False - master_config.grpo["use_dynamic_sampling"] = False + master_config.grpo["use_dynamic_sampling"] = True grpo_mod.grpo_train( mock_grpo_components["policy"], @@ -1913,6 +1940,12 @@ def fake_batched_message_log_to_flat_message(*_args, **_kwargs): assert train_metrics["min_seq_mult_prob_error_after_mask"] == 1.0 assert train_metrics["num_masked_seqs_by_logprob_error"] == 2 assert train_metrics["masked_correct_pct"] == 0.5 + assert dynamic_sampling_calls == 2 + assert any( + call.args[0] == {"delta/changed_pct": 4.0} + and call.kwargs.get("prefix") == "refit" + for call in mock_grpo_components["logger"].log_metrics.call_args_list + ) def test_grpo_train_shutdown_on_epoch_completion(mock_grpo_components, tmp_path): diff --git a/tests/unit/models/generation/test_vllm_backend.py b/tests/unit/models/generation/test_vllm_backend.py index a63242a9f0f..0d29ed7d60e 100644 --- a/tests/unit/models/generation/test_vllm_backend.py +++ b/tests/unit/models/generation/test_vllm_backend.py @@ -18,7 +18,7 @@ import contextlib import json -from types import SimpleNamespace +from types import MethodType, SimpleNamespace from typing import Any from unittest.mock import MagicMock @@ -26,6 +26,8 @@ import torch from safetensors.torch import save_file +from nemo_rl.utils.weight_transfer_sparse_codec import encode_sparse_infos + def _make_collective_update_extension(backend): ext = backend.VllmInternalWorkerExtension.__new__( @@ -98,7 +100,6 @@ def _make_sparse_delta_extension( parameter_name: str, target: torch.Tensor, module: object, - module_name: str | None = None, ) -> Any: from nemo_rl.models.generation.vllm.vllm_backend import ( VllmInternalWorkerExtension, @@ -107,10 +108,12 @@ def _make_sparse_delta_extension( ext = VllmInternalWorkerExtension.__new__(VllmInternalWorkerExtension) ext.rank = 1 ext._direct_sparse_delta_targets = {parameter_name: target} - ext._direct_sparse_delta_modules = { - module_name or parameter_name.rsplit(".", 1)[0]: module - } + ext.model_runner = SimpleNamespace( + model=SimpleNamespace(get_submodule=lambda _name: module) + ) ext._direct_sparse_delta_plan_cache = {} + ext._direct_sparse_delta_verification = None + ext._direct_sparse_delta_verification_candidates = 0 return ext @@ -145,10 +148,10 @@ def test_serialized_sparse_payload_batch_preserves_order(tmp_path) -> None: paths = [tmp_path / f"{index}.pt" for index in range(3)] for path, payload in zip(paths, payloads, strict=True): torch.save(payload, path) - applied: list[tuple[Any, bool]] = [] + applied: list[Any] = [] - def apply(payload: Any, *, synchronize: bool) -> dict[str, Any]: - applied.append((payload, synchronize)) + def apply(payload: Any) -> dict[str, Any]: + applied.append(payload) return { "ok": True, "receiver_sparse_apply_s": 2.0, @@ -156,15 +159,14 @@ def apply(payload: Any, *, synchronize: bool) -> dict[str, Any]: ext._apply_sparse_request = apply result = ext.update_weights_from_sparse_payload_files( - *(str(path) for path in paths), synchronize=False + *(str(path) for path in paths) ) - assert [item[0][2]["index"] for item in applied] == [0, 1, 2] + assert [item[2]["index"] for item in applied] == [0, 1, 2] assert all( - torch.equal(item[0][1], payload[1]) + torch.equal(item[1], payload[1]) for item, payload in zip(applied, payloads, strict=True) ) - assert all(not synchronize for _, synchronize in applied) assert result["receiver_deserialize_s"] >= 0.0 assert result["receiver_sparse_apply_s"] == 6.0 @@ -189,6 +191,30 @@ def test_direct_sparse_delta_placement() -> None: ) _assert_sparse_plan(ext, plan, [0, 1, 2, 3], [8, 9, 10, 11], [0.0, 1.0, 2.0, 3.0]) + merged_name = "model.layers.0.mlp.gate_up_proj.weight" + merged_target = _attach_tensor_attrs(torch.zeros(8, 2), output_dim=0) + ext = _make_sparse_delta_extension( + merged_name, + merged_target, + SimpleNamespace(tp_rank=1, tp_size=2, output_sizes=(8, 8)), + ) + for projection, expected_locations in ( + ("gate", [0, 1, 6, 7]), + ("up", [8, 9, 14, 15]), + ): + source_name = f"model.layers.0.mlp.{projection}_proj.weight" + plan = ext._direct_sparse_delta_target_plan( + {"name": source_name, "shape": (8, 2)}, + {merged_name: merged_target}, + ) + _assert_sparse_plan( + ext, + plan, + [6, 7, 8, 9, 14, 15], + expected_locations, + [2.0, 3.0, 4.0, 5.0], + ) + expert_name = "model.layers.0.mlp.experts.w13_weight" expert_target = torch.zeros(2, 4, 2) expert_module = SimpleNamespace( @@ -203,33 +229,20 @@ def test_direct_sparse_delta_placement() -> None: expert_target, expert_module, ) - expert_source = "model.layers.0.mlp.experts.3.gate_proj.weight" - plan = ext._direct_sparse_delta_expert_plan( - {"name": expert_source, "shape": (8, 2)}, - expert_source, - {expert_name: expert_target}, - ) - _assert_sparse_plan( - ext, - plan, - [6, 7, 8, 9, 14, 15], - [8, 9, 14, 15], - [2.0, 3.0, 4.0, 5.0], - ) - - expert_source = "model.layers.0.mlp.experts.3.up_proj.weight" - plan = ext._direct_sparse_delta_expert_plan( - {"name": expert_source, "shape": (8, 2)}, - expert_source, - {expert_name: expert_target}, - ) - _assert_sparse_plan( - ext, - plan, - [6, 7, 8, 9, 14, 15], - [8, 9, 14, 15], - [2.0, 3.0, 4.0, 5.0], - ) + for projection in ("gate_proj", "up_proj"): + expert_source = f"model.layers.0.mlp.experts.3.{projection}.weight" + plan = ext._direct_sparse_delta_expert_plan( + {"name": expert_source, "shape": (8, 2)}, + expert_source, + {expert_name: expert_target}, + ) + _assert_sparse_plan( + ext, + plan, + [6, 7, 8, 9, 14, 15], + [8, 9, 14, 15], + [2.0, 3.0, 4.0, 5.0], + ) w2_target = torch.zeros(2, 2, 4) ext = _make_sparse_delta_extension(expert_name, w2_target, expert_module) @@ -242,81 +255,120 @@ def test_direct_sparse_delta_placement() -> None: _assert_sparse_plan(ext, plan, [3, 4, 7, 11, 15], [8, 11, 15], [1.0, 2.0, 4.0]) mamba_name = "model.layers.0.mixer.in_proj.weight" - mamba_target = torch.zeros(16, 1, 2) - ext = _make_sparse_delta_extension( - mamba_name, - mamba_target, - SimpleNamespace( - tp_size=2, - intermediate_size=8, - groups_ssm_state_size=6, - num_heads=4, + for target_shape, groups, source_locations, expected_locations, values in ( + ((16, 1, 2), 6, [0, 8, 25, 38, 40, 55], [0, 9, 22, 31], [1, 2, 4, 5]), + ( + (14, 2), + 4, + [0, 8, 24, 36, 44, 52, 55], + [0, 8, 16, 20, 24, 27], + [1, 2, 3, 4, 5, 6], ), - "model.layers.0.mixer", - ) - plan = ext._direct_sparse_delta_mamba2_plan( - {"name": mamba_name, "shape": (28, 2)}, - mamba_name, - {mamba_name: mamba_target}, - ) - _assert_sparse_plan( - ext, plan, [0, 8, 25, 38, 40, 55], [0, 9, 22, 31], [1.0, 2.0, 4.0, 5.0] + ): + target = _attach_tensor_attrs( + torch.zeros(target_shape), + weight_loader=MethodType(lambda _owner: None, SimpleNamespace()), + ) + ext = _make_sparse_delta_extension( + mamba_name, + target, + SimpleNamespace( + tp_size=2, + intermediate_size=8, + groups_ssm_state_size=groups, + num_heads=4, + ), + ) + plan = ext._direct_sparse_delta_mamba2_plan( + {"name": mamba_name, "shape": (28, 2)}, + mamba_name, + {mamba_name: target}, + ) + _assert_sparse_plan(ext, plan, source_locations, expected_locations, values) + + for attrs, source_shape, source_locations, expected_locations, values in ( + ( + {"output_dim": 0}, + (6, 2), + [0, 1, 6, 7, 10, 11], + [0, 1, 4, 5], + [2, 3, 4, 5], + ), + ( + {"output_dim": 0, "input_dim": 1}, + (3, 4), + [0, 1, 2, 3, 6, 7, 10, 11], + [0, 1, 2, 3, 4, 5], + [2, 3, 4, 5, 6, 7], + ), + ): + target = _attach_tensor_attrs(torch.zeros(3, 2), **attrs, tp_size=2, tp_rank=1) + ext = _make_sparse_delta_extension("down_proj.weight", target, object()) + plan = ext._direct_sparse_delta_shard_plan( + {"name": "down_proj.weight", "shape": source_shape}, target + ) + _assert_sparse_plan(ext, plan, source_locations, expected_locations, values) + + +@pytest.mark.vllm +@pytest.mark.parametrize( + ("initial", "expected_delta", "exact_mismatches", "mismatches"), + [ + (200.0, 4.0, 0, 0), + (2.0, 4.0000005, 1, 0), + (2.0, 5.0, 1, 1), + ], +) +def test_sparse_delta_sample_verification_only_compares_applied_delta( + monkeypatch, + initial: float, + expected_delta: float, + exact_mismatches: int, + mismatches: int, +) -> None: + from nemo_rl.models.generation.vllm.vllm_backend import ( + VllmInternalWorkerExtension, ) - mamba_target = torch.zeros(14, 2) - ext = _make_sparse_delta_extension( - mamba_name, - mamba_target, - SimpleNamespace( - tp_size=2, - intermediate_size=8, - groups_ssm_state_size=4, - num_heads=4, + monkeypatch.setattr(torch.cuda, "is_available", lambda: False) + monkeypatch.setattr( + "nemo_rl.models.generation.vllm.quantization.fp8.is_fp8_model", + lambda _config: False, + ) + target = torch.tensor([1.0, initial, 3.0, initial]) + ext = VllmInternalWorkerExtension.__new__(VllmInternalWorkerExtension) + ext.model_runner = SimpleNamespace( + model=SimpleNamespace(), + vllm_config=SimpleNamespace( + model_config=SimpleNamespace(architectures=[]), ), - "model.layers.0.mixer", ) - plan = ext._direct_sparse_delta_mamba2_plan( - {"name": mamba_name, "shape": (28, 2)}, - mamba_name, - {mamba_name: mamba_target}, + ext._direct_sparse_delta_targets = {"weight": target} + ext._direct_sparse_delta_plan_cache = { + "weight": ext._make_sparse_delta_target_plan(target, (4,)) + } + ext._direct_sparse_delta_verification = None + ext._direct_sparse_delta_verification_candidates = 0 + payload = encode_sparse_infos( + [("weight", target, torch.tensor([1, 3]), torch.tensor([4.0, 4.0]))], + empty_dtype=target.dtype, ) - _assert_sparse_plan( - ext, - plan, - [0, 8, 24, 36, 44, 52, 55], - [0, 8, 16, 20, 24, 27], - [1, 2, 3, 4, 5, 6], + metadata = payload[2] + metadata[0].update( + verification_locations=[1, 3], + verification_deltas=[expected_delta, expected_delta], ) - shard_target = _attach_tensor_attrs( - torch.zeros(3, 2), output_dim=0, tp_size=2, tp_rank=1 - ) - ext = _make_sparse_delta_extension("down_proj.weight", shard_target, object()) - plan = ext._direct_sparse_delta_shard_plan( - {"name": "down_proj.weight", "shape": (6, 2)}, shard_target - ) - _assert_sparse_plan( - ext, - plan, - [0, 1, 6, 7, 10, 11], - [0, 1, 4, 5], - [2.0, 3.0, 4.0, 5.0], - ) + ext._apply_sparse_weight_deltas(payload[:2], metadata) + result = ext.finish_sparse_delta_refit() - shared_target = _attach_tensor_attrs( - torch.zeros(3, 2), output_dim=0, input_dim=1, tp_size=2, tp_rank=1 - ) - ext = _make_sparse_delta_extension("down_proj.weight", shared_target, object()) - plan = ext._direct_sparse_delta_shard_plan( - {"name": "down_proj.weight", "shape": (3, 4)}, shared_target - ) - _assert_sparse_plan( - ext, - plan, - [0, 1, 2, 3, 6, 7, 10, 11], - [0, 1, 2, 3, 4, 5], - [2.0, 3.0, 4.0, 5.0, 6.0, 7.0], - ) + assert torch.equal(target, torch.tensor([1.0, initial + 4.0, 3.0, initial + 4.0])) + assert result["verification_candidates"] == 2 + assert result["verification_samples"] == 2 + assert result["verification_exact_mismatches"] == 2 * exact_mismatches + assert result["verification_mismatches"] == 2 * mismatches + rounded_difference = float((torch.tensor(expected_delta) - 4).abs()) + assert result["verification_max_abs"] == rounded_difference @pytest.mark.vllm @@ -325,8 +377,10 @@ def test_update_weights_from_collective_processes_weights_after_loading(monkeypa call_order = [] process_calls = [] + current_configs = [] def process_weights_after_loading(model, model_config, device): + assert current_configs == [ext.model_runner.vllm_config] call_order.append("process") process_calls.append((model, model_config, device)) @@ -372,6 +426,7 @@ def packed_broadcast_consumer(iterator, group, src, post_unpack_func): assert ext.update_weights_from_collective() is True + assert not current_configs assert process_calls == [(ext.model_runner.model, ext.model_config, ext.device)] assert call_order == [ "broadcast", diff --git a/tests/unit/models/generation/test_vllm_generation.py b/tests/unit/models/generation/test_vllm_generation.py index de870b6325d..5a6e5dc579c 100644 --- a/tests/unit/models/generation/test_vllm_generation.py +++ b/tests/unit/models/generation/test_vllm_generation.py @@ -12,18 +12,21 @@ # See the License for the specific language governing permissions and # limitations under the License. +import asyncio import importlib.util import json import os import sys import threading +import time import types -from collections import deque -from concurrent.futures import ThreadPoolExecutor +from collections.abc import Iterator +from concurrent.futures import Future, ThreadPoolExecutor +from contextlib import contextmanager from copy import deepcopy from pathlib import Path from typing import Any -from unittest.mock import MagicMock +from unittest.mock import AsyncMock, MagicMock, call import pytest import ray @@ -165,56 +168,143 @@ def test_resolve_enable_prefix_caching_uses_cuda_capability_for_auto(monkeypatch assert _resolve_enable_prefix_caching({}) is False -def test_sparse_refit_queue_batches_payloads_in_fifo_order() -> None: +def test_ray_owner_destructors_skip_shutdown_during_interpreter_finalization( + monkeypatch, +): + monkeypatch.setattr(sys, "is_finalizing", lambda: True) + + generation = VllmGeneration.__new__(VllmGeneration) + generation.shutdown = MagicMock() + generation.__del__() + generation.shutdown.assert_not_called() + + policy = Policy.__new__(Policy) + policy.worker_group = MagicMock() + policy.__del__() + policy.worker_group.shutdown.assert_not_called() + + cluster = RayVirtualCluster.__new__(RayVirtualCluster) + cluster.shutdown = MagicMock() + cluster.__del__() + cluster.shutdown.assert_not_called() + + +@contextmanager +def _sparse_refit_worker( + *, batch_size: int = 2, futures: list[Future[dict[str, Any]]] | None = None +) -> Iterator[BaseVllmGenerationWorker]: worker = BaseVllmGenerationWorker.__new__(BaseVllmGenerationWorker) - worker._refit_apply_queue_lock = threading.Lock() + worker._refit_apply_queue_condition = threading.Condition() worker._refit_apply_executor = ThreadPoolExecutor(max_workers=1) - worker._refit_apply_futures = deque() + worker._refit_apply_futures = list(futures or []) worker._refit_apply_pending_payloads = [] - worker._refit_apply_payload_count = 0 - worker._refit_apply_batch_count = 0 + worker._refit_seen_payloads = {} worker._refit_apply_queue_depth = 2 - worker._refit_apply_batch_size = 3 + worker._refit_apply_batch_size = batch_size worker.llm = MagicMock() - applied: list[tuple[tuple[bytes, ...], bool]] = [] + worker.llm.collective_rpc.return_value = [{"ok": True}] + for future in worker._refit_apply_futures: + future.add_done_callback(worker._notify_refit_apply_waiters) + try: + yield worker + finally: + worker._refit_apply_executor.shutdown(wait=True) + - def apply(payloads: tuple[bytes, ...], synchronize: bool) -> dict[str, Any]: - applied.append((payloads, synchronize)) +def test_sparse_refit_queue_batches_payloads_in_fifo_order() -> None: + applied: list[tuple[bytes, ...]] = [] + + def apply(payloads: tuple[bytes, ...]) -> dict[str, Any]: + applied.append(payloads) return { "ok": True, "payloads": len(payloads), "receiver_total_s": float(len(payloads)), } - worker.update_weights_from_serialized_sparse_payloads = apply - try: + with _sparse_refit_worker(batch_size=3) as worker: + worker.update_weights_from_serialized_sparse_payloads = apply responses = [ - worker._enqueue_sparse_payload_apply(payload) - for payload in (b"0", b"1", b"2", b"3", b"4") + worker._enqueue_sparse_payload_apply( + payload, ("transfer", 0, index), str(index) + ) + for index, payload in enumerate((b"0", b"1", b"2", b"3", b"4")) ] response = worker._flush_queued_sparse_payloads() responses.append(response) - finally: - worker._refit_apply_executor.shutdown(wait=True) assert applied == [ - ((b"0", b"1", b"2"), False), - ((b"3", b"4"), False), + (b"0", b"1", b"2"), + (b"3", b"4"), ] assert response["payloads"] == 5 assert response["batches"] == 2 assert sum(result.get("receiver_total_s", 0.0) for result in responses) == 5.0 - worker.llm.collective_rpc.assert_called_once_with("synchronize_device", args=()) + worker.llm.collective_rpc.assert_called_once_with( + "finish_sparse_delta_refit", args=() + ) + + +def test_sparse_refit_queue_deduplicates_transactional_payloads() -> None: + key = ("transfer", 0, 1) + with _sparse_refit_worker() as worker: + worker.update_weights_from_serialized_sparse_payloads = MagicMock( + return_value={"ok": True, "payloads": 1} + ) + assert worker._enqueue_sparse_payload_apply(b"payload", key, "checksum")["ok"] + duplicate = worker._enqueue_sparse_payload_apply(b"payload", key, "checksum") + assert duplicate == {"ok": True, "payloads": 0, "duplicate": True} + with pytest.raises(ValueError, match="reused with different data"): + worker._enqueue_sparse_payload_apply(b"other", key, "different") + response = worker._flush_queued_sparse_payloads() + + assert response["payloads"] == 1 + assert worker._refit_seen_payloads == {} + + +def test_sparse_refit_queue_does_not_deduplicate_failed_enqueue() -> None: + failed = Future() + failed.set_exception(RuntimeError("prior apply failed")) + with _sparse_refit_worker(futures=[failed]) as worker: + with pytest.raises(RuntimeError, match="prior apply failed"): + worker._enqueue_sparse_payload_apply( + b"payload", ("transfer", 0, 1), "checksum" + ) + + assert worker._refit_seen_payloads == {} + assert worker._refit_apply_pending_payloads == [] + + +def test_sparse_refit_queue_releases_condition_while_backpressured() -> None: + first: Future[dict[str, Any]] = Future() + second: Future[dict[str, Any]] = Future() + started = threading.Event() + + with _sparse_refit_worker(futures=[first, second]) as worker: + with ThreadPoolExecutor(max_workers=1) as callers: + call = callers.submit( + lambda: ( + started.set(), + worker._enqueue_sparse_payload_apply( + b"payload", ("transfer", 0, 1), "checksum" + ), + )[1] + ) + assert started.wait(timeout=1.0) + time.sleep(0.05) + assert worker._refit_apply_queue_condition.acquire(timeout=1.0) + worker._refit_apply_queue_condition.release() + first.set_result({"ok": True, "payloads": 1}) + assert call.result(timeout=1.0)["ok"] def test_sparse_refit_batch_uses_one_collective_rpc(tmp_path: Path) -> None: worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) staged_payloads: list[bytes] = [] - def collective_rpc(method, args, kwargs): + def collective_rpc(method, args): assert method == "update_weights_from_sparse_payload_files" staged_payloads.extend(Path(path).read_bytes() for path in args) - assert kwargs == {"synchronize": False} return [{"ok": True, "receiver_total_s": 1.0}] worker.llm = MagicMock(collective_rpc=MagicMock(side_effect=collective_rpc)) @@ -222,9 +312,7 @@ def collective_rpc(method, args, kwargs): worker._refit_batch_staging_dir = str(tmp_path) payloads = (b"0", b"1", b"2") - response = worker.update_weights_from_serialized_sparse_payloads( - payloads, synchronize=False - ) + response = worker.update_weights_from_serialized_sparse_payloads(payloads) assert staged_payloads == list(payloads) assert not list(tmp_path.iterdir()) @@ -232,6 +320,33 @@ def collective_rpc(method, args, kwargs): assert response == {"ok": True, "receiver_total_s": 1.0, "payloads": 3} +def test_sparse_refit_batch_drains_workers_before_error_cleanup(tmp_path: Path) -> None: + worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) + staged_paths: tuple[str, ...] = () + + def collective_rpc(method, args): + nonlocal staged_paths + if method == "update_weights_from_sparse_payload_files": + staged_paths = args + raise RuntimeError("apply failed") + assert method == "synchronize_device" + assert all(Path(path).is_file() for path in staged_paths) + return [True] + + worker.llm = MagicMock(collective_rpc=MagicMock(side_effect=collective_rpc)) + worker._refit_workers_share_node = True + worker._refit_batch_staging_dir = str(tmp_path) + + with pytest.raises(RuntimeError, match="apply failed"): + worker.update_weights_from_serialized_sparse_payloads((b"0", b"1")) + + assert [call.args[0] for call in worker.llm.collective_rpc.call_args_list] == [ + "update_weights_from_sparse_payload_files", + "synchronize_device", + ] + assert not list(tmp_path.iterdir()) + + def test_sparse_refit_batch_falls_back_across_nodes() -> None: worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) worker._refit_workers_share_node = False @@ -241,21 +356,87 @@ def test_sparse_refit_batch_falls_back_across_nodes() -> None: response = worker.update_weights_from_serialized_sparse_payloads((b"0", b"1", b"2")) - calls = worker.llm.collective_rpc.call_args_list - assert [call.args[0] for call in calls] == [ - "update_weights_from_serialized_sparse_payload", - "update_weights_from_serialized_sparse_payload", - "update_weights_from_serialized_sparse_payload", - "synchronize_device", - ] - assert [call.kwargs["args"] for call in calls[:3]] == [ - (b"0", False), - (b"1", False), - (b"2", False), + assert worker.llm.collective_rpc.call_args_list == [ + call("update_weights_from_serialized_sparse_payload", args=(payload,)) + for payload in (b"0", b"1", b"2") ] assert response == {"ok": True, "receiver_total_s": 3.0, "payloads": 3} +@pytest.mark.asyncio +async def test_async_sparse_refit_batch_bridges_to_async_collective( + tmp_path: Path, +) -> None: + worker = VllmAsyncGenerationWorkerImpl.__new__(VllmAsyncGenerationWorkerImpl) + staged_payloads: list[bytes] = [] + + class AsyncLlm: + async def collective_rpc( + self, method: str, args: tuple[Any, ...] + ) -> list[dict[str, Any]]: + assert method == "update_weights_from_sparse_payload_files" + staged_payloads.extend(Path(path).read_bytes() for path in args) + return [{"ok": True, "receiver_total_s": 1.0}] + + worker.llm = AsyncLlm() + worker._refit_async_loop = asyncio.get_running_loop() + worker._refit_workers_share_node = True + worker._refit_batch_staging_dir = str(tmp_path) + + response = await asyncio.to_thread( + worker.update_weights_from_serialized_sparse_payloads, + (b"0", b"1"), + ) + + assert staged_payloads == [b"0", b"1"] + assert response == {"ok": True, "receiver_total_s": 1.0, "payloads": 2} + + +@pytest.mark.asyncio +async def test_async_sparse_refit_post_init_records_worker_locality() -> None: + worker = VllmAsyncGenerationWorkerImpl.__new__(VllmAsyncGenerationWorkerImpl) + worker.cfg = {"refit_transport": "vllm_zmq_sparse"} + worker._mtp_load_from_disk = False + worker.report_device_id_async = AsyncMock(return_value=["0"]) + worker.llm = MagicMock() + worker.llm.collective_rpc = AsyncMock(return_value=["node-0", "node-0"]) + + await worker.post_init_async() + + assert worker.vllm_device_ids == ["0"] + assert worker._refit_workers_share_node is True + worker.llm.collective_rpc.assert_awaited_once_with("report_node_hostname", args=()) + + +def test_async_sparse_refit_exposes_zmq_relay(monkeypatch) -> None: + server = MagicMock() + server_type = MagicMock(return_value=server) + monkeypatch.setattr( + "nemo_rl.models.generation.vllm.vllm_worker.ZmqSparseRefitServer", + server_type, + ) + monkeypatch.setattr( + "nemo_rl.models.generation.vllm.vllm_worker._get_free_port_local", + lambda *_args: 12345, + ) + monkeypatch.setattr( + "nemo_rl.models.generation.vllm.vllm_worker._get_node_ip_local", + lambda: "10.0.0.1", + ) + worker = VllmAsyncGenerationWorkerImpl.__new__(VllmAsyncGenerationWorkerImpl) + worker.cfg = {"vllm_cfg": {"zmq_refit_server_port": None}} + worker._zmq_refit_server = None + + assert worker.start_zmq_sparse_refit_relay(["http://receiver"]) == ( + "tcp://10.0.0.1:12345" + ) + server.start.assert_called_once_with() + + worker.stop_zmq_sparse_refit_relay() + server.close.assert_called_once_with() + assert worker._zmq_refit_server is None + + basic_lora_test_config: LoRAConfig = { "enabled": False, "target_modules": [], @@ -544,9 +725,10 @@ def test_configure_generation_config_uses_real_startup_weights_without_draft_ref assert configured["vllm_cfg"]["load_format"] == "auto" -def test_configure_generation_config_uses_real_s3_delta_baseline(): +@pytest.mark.parametrize("transport", ["vllm_s3_sparse", "vllm_zmq_sparse"]) +def test_configure_generation_config_uses_real_delta_baseline(transport: str): vllm_config = deepcopy(basic_vllm_test_config) - vllm_config["refit_transport"] = "vllm_s3_sparse" + vllm_config["refit_transport"] = transport configured = configure_generation_config( vllm_config, MagicMock(pad_token_id=0, eos_token_id=1) diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index 793cc52bd90..7e284ff2f85 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 Optional +from unittest.mock import Mock import numpy as np import pytest @@ -75,6 +76,67 @@ def test_megatron_prepare_for_training_restores_optimizer(): assert restored_devices == ["cuda"] +def test_remote_sparse_stream_drains_cuda_before_return(monkeypatch): + import nemo_rl.models.policy.workers.megatron_policy_worker as worker_module + + worker = object.__new__(worker_module.MegatronPolicyWorkerImpl) + worker.delta_weight_transfer_tracker = object() + worker._iter_params_with_optional_kv_scales = lambda: iter(()) + result = {"payloads": 1, "changed_elements": 2, "total_elements": 3} + events = [] + + def stream(*_args, **_kwargs): + events.append("stream") + return result + + monkeypatch.setattr(worker_module, "stream_sparse_delta_payloads_via_zmq", stream) + monkeypatch.setattr(torch.cuda, "is_available", lambda: True) + monkeypatch.setattr(torch.cuda, "synchronize", lambda: events.append("sync")) + monkeypatch.setattr(torch.cuda.nvtx, "range_push", lambda _name: None) + monkeypatch.setattr(torch.cuda.nvtx, "range_pop", lambda: None) + + actual = worker_module.MegatronPolicyWorkerImpl.stream_remote_sparse_weights( + worker, + "zmq", + ["tcp://receiver:5555"], + transfer_id="transfer", + api_key_env_var=None, + timeout_s=1.0, + shard_rank=0, + shard_count=1, + ) + + assert actual is result + assert events == ["stream", "sync"] + + +def test_refit_param_info_uses_mapping_pp_broadcast(): + from nemo_rl.models.policy.workers.megatron_policy_worker import ( + MegatronPolicyWorkerImpl, + ) + + mapping = SimpleNamespace( + tp_size=2, + ep_size=4, + is_expert=True, + broadcast_obj_from_pp_rank=Mock(return_value=192), + ) + worker = object.__new__(MegatronPolicyWorkerImpl) + worker.dtype = torch.bfloat16 + worker.refit_conversion_tasks = [ + SimpleNamespace( + param_name="decoder.layers.0.mlp.weight", + param_weight=torch.empty(3, 4, dtype=torch.float16), + mapping=mapping, + ) + ] + + assert worker._calculate_refit_param_info() == [ + ("decoder.layers.0.mlp.weight", 192) + ] + mapping.broadcast_obj_from_pp_rank.assert_called_once_with(192) + + def test_set_moe_grad_scale_func_sets_and_clears_on_model_config(): """_set_moe_grad_scale_func should set/clear moe_grad_scale_func on the config.""" from nemo_rl.models.policy.workers.megatron_policy_worker import ( diff --git a/tests/unit/reference_configs/grpo_math_1B.yaml b/tests/unit/reference_configs/grpo_math_1B.yaml index 48861f83257..217556bfc2a 100644 --- a/tests/unit/reference_configs/grpo_math_1B.yaml +++ b/tests/unit/reference_configs/grpo_math_1B.yaml @@ -332,8 +332,8 @@ policy: top_k: null stop_token_ids: null stop_strings: null - refit_transport: null # Set to "vllm_s3_sparse" to use S3 sparse-delta refit. - delta_compression: null # S3 sparse-delta refit config; null uses the existing refit path. + refit_transport: null # Set to "vllm_s3_sparse" or "vllm_zmq_sparse" for remote sparse-delta refit. + delta_compression: null # Remote sparse-delta config; null uses the existing refit path. mcore_generation_config: async_engine: false max_model_len: ${policy.max_total_sequence_length} @@ -371,9 +371,9 @@ policy: num_first_layers_in_bf16: 0 enable_vllm_metrics_logger: true # Set to true to enable vLLM internal metrics logger, turn off for better performance vllm_metrics_logger_interval: 0.5 # Interval in seconds to collect vLLM logger metrics - expose_http_refit_server: false # Start the internal sparse-delta refit endpoint on vLLM workers. http_refit_api_key_env_var: null # Optional env var containing the internal refit API key. http_refit_server_port: null # Optional fixed port for Kubernetes targetPorts. + zmq_refit_server_port: null # Optional fixed ZeroMQ relay port for Kubernetes targetPorts. vllm_kwargs: {} colocated: # true: generation shares training GPUs diff --git a/tests/unit/utils/test_weight_transfer_remote_sparse.py b/tests/unit/utils/test_weight_transfer_remote_sparse.py new file mode 100644 index 00000000000..906ae092d4f --- /dev/null +++ b/tests/unit/utils/test_weight_transfer_remote_sparse.py @@ -0,0 +1,306 @@ +# 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 json +import threading +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from types import SimpleNamespace + +import pytest +import torch +import zstandard + +from nemo_rl.utils import weight_transfer_remote_sparse +from nemo_rl.utils.weight_transfer_remote_sparse import download_s3_refit_payload +from nemo_rl.utils.weight_transfer_sparse_codec import ( + DeltaCompressionTracker, + encode_sparse_infos, + sparse_locations_for_item, +) +from nemo_rl.utils.weight_transfer_zmq import ( + G_VLLM_REFIT_CHECKSUM_HEADER, + G_VLLM_REFIT_PAYLOAD_HEADER, + G_VLLM_REFIT_PRODUCER_HEADER, + G_VLLM_REFIT_TRANSFER_HEADER, + ZmqSparseRefitClient, + ZmqSparseRefitServer, + sparse_payload_checksum, +) + + +def test_delta_tracker_commits_only_successful_syncs(monkeypatch) -> None: + monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") + tracker = DeltaCompressionTracker( + {"dtype": "bf16", "sparse_bucket_size_bytes": 1024} + ) + tensor = torch.tensor([1.0, 2.0, 3.0]) + tracker.snapshot_baseline([("weight", tensor)]) + tensor[1] += 4 + + assert tracker.prepare_sparse_delta_payload([("weight", tensor)])[0][2] + tracker.on_sync_failed() + assert tracker.prepare_sparse_delta_payload([("weight", tensor)])[0][2] + tracker.on_sync_succeeded() + assert not tracker.prepare_sparse_delta_payload([("weight", tensor)])[0][2] + + +def test_delta_tracker_emits_bounded_delta_samples(monkeypatch) -> None: + monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") + monkeypatch.setenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "2") + tracker = DeltaCompressionTracker( + {"dtype": "bf16", "sparse_bucket_size_bytes": 1024} + ) + tensor = torch.tensor([1.0, 2.0, 3.0, 4.0]) + tracker.snapshot_baseline([("weight", tensor)]) + tensor[[1, 3]] += 1 + + (_, _, metadata), changed, total = tracker.prepare_sparse_delta_payload( + [("weight", tensor)] + ) + + assert metadata[0]["verification_locations"] == [1, 3] + assert metadata[0]["verification_deltas"] == [1.0, 1.0] + assert (changed, total) == (2, 4) + + +def test_delta_tracker_commits_quantized_receiver_baseline(monkeypatch) -> None: + monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") + tracker = DeltaCompressionTracker( + {"dtype": "bf16", "sparse_bucket_size_bytes": 1024} + ) + tensor = torch.tensor([1.0]) + tracker.snapshot_baseline([("weight", tensor)]) + tensor.add_(0.001) + + (_, deltas, _), _, _ = tracker.prepare_sparse_delta_payload([("weight", tensor)]) + expected = torch.tensor([1.0]) + deltas.float() + tracker.on_sync_succeeded() + tracker.prepare_sparse_delta_payload([("weight", tensor)]) + + assert torch.equal(tracker.baseline["weight"], expected) + assert not torch.equal(expected, tensor) + + +def test_sparse_index_encoding_preserves_uint64_locations() -> None: + locations = torch.tensor([0, 2**32 + 5]) + packed, _, metadata = encode_sparse_infos( + [("weight", torch.empty(2), locations, torch.ones(2))], + empty_dtype=torch.float32, + ) + + decoded = sparse_locations_for_item(metadata[0], packed, device="cpu") + assert torch.equal(decoded, locations) + + +def test_s3_download_verifies_checksum(monkeypatch) -> None: + compressed = zstandard.ZstdCompressor().compress(b"payload") + monkeypatch.setattr( + weight_transfer_remote_sparse, + "_get_manifest_s3_store", + lambda *_args: SimpleNamespace(get=lambda _key: compressed), + ) + manifest = { + "bucket": "bucket", + "region": "region", + "key": "key", + "checksum": sparse_payload_checksum(compressed), + } + + assert download_s3_refit_payload(manifest) == b"payload" + manifest["checksum"] = "0" * 32 + with pytest.raises(ValueError, match="checksum mismatch"): + download_s3_refit_payload(manifest) + + +def test_refit_http_session_does_not_retry_application_errors() -> None: + retry = ( + weight_transfer_remote_sparse.refit_http_session() + .get_adapter("http://") + .max_retries + ) + + assert 500 not in retry.status_forcelist + assert {502, 503, 504} <= set(retry.status_forcelist) + + +def test_sparse_export_finishes_before_blocked_transfers(monkeypatch) -> None: + monkeypatch.setenv("NRL_REFIT_ZMQ_ENCODE_WORKERS", "1") + monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "1") + exported = threading.Event() + release_transfers = threading.Event() + result = [] + + class Tracker: + sparse_bucket_size_bytes = 1 + + @staticmethod + def prepare_sparse_delta_payload(chunk): + return (chunk, torch.ones(1), [1]), 1, 1 + + def tensors(): + for index in range(4): + yield f"weight-{index}", torch.ones(1) + exported.set() + + def send_payload(_body, _payload_index): + assert release_transfers.wait(timeout=5.0) + return {"receiver": {}} + + def run(): + result.append( + weight_transfer_remote_sparse.stream_sparse_delta_payloads( + tensors(), + delta_tracker=Tracker(), + transport="zmq", + send_payload=send_payload, + transfer_workers=1, + shard_rank=0, + shard_count=1, + ) + ) + + thread = threading.Thread(target=run) + thread.start() + try: + assert exported.wait(timeout=2.0) + finally: + release_transfers.set() + thread.join(timeout=5.0) + + assert not thread.is_alive() + assert result == [{"payloads": 4, "changed_elements": 4, "total_elements": 4}] + + +def test_sparse_export_finishes_before_transfer_error(monkeypatch) -> None: + monkeypatch.setenv("NRL_REFIT_ZMQ_ENCODE_WORKERS", "1") + monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "1") + exported = [] + + class Tracker: + sparse_bucket_size_bytes = 1 + + @staticmethod + def prepare_sparse_delta_payload(chunk): + return (chunk, torch.ones(1), [1]), 1, 1 + + def tensors(): + for index in range(4): + exported.append(index) + yield f"weight-{index}", torch.ones(1) + + def fail_transfer(_body, _payload_index): + raise RuntimeError("transfer failed") + + with pytest.raises(RuntimeError, match="transfer failed"): + weight_transfer_remote_sparse.stream_sparse_delta_payloads( + tensors(), + delta_tracker=Tracker(), + transport="zmq", + send_payload=fail_transfer, + transfer_workers=1, + shard_rank=0, + shard_count=1, + ) + + assert exported == list(range(4)) + + +def _receiver_server(received): + class Handler(BaseHTTPRequestHandler): + def do_POST(self): + body = self.rfile.read(int(self.headers["content-length"])) + received.append( + ({key.lower(): value for key, value in self.headers.items()}, body) + ) + response = json.dumps({"ok": True, "receiver_total_s": 0.25}).encode() + self.send_response(200) + self.send_header("content-type", "application/json") + self.send_header("content-length", str(len(response))) + self.end_headers() + self.wfile.write(response) + + def log_message(self, *args): + pass + + server = ThreadingHTTPServer(("127.0.0.1", 0), Handler) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + return server, thread + + +def _send_zmq_payload( + client: ZmqSparseRefitClient, + payload_id: int, + body: bytes, + checksum: str | None = None, +) -> dict[str, object]: + return client.send_payload( + transfer_id="transfer-a", + payload_id=payload_id, + checksum=checksum or sparse_payload_checksum(body), + body=body, + ) + + +def test_zmq_sparse_refit_relay_fans_out_and_rejects_corruption(monkeypatch) -> None: + monkeypatch.setenv("NRL_TEST_REFIT_KEY", "secret") + received = [[], []] + receivers = [_receiver_server(items) for items in received] + urls = [f"http://127.0.0.1:{server.server_port}" for server, _ in receivers] + relay = ZmqSparseRefitServer( + urls, + bind_address="tcp://127.0.0.1:*", + api_key_env_var="NRL_TEST_REFIT_KEY", + timeout_s=5.0, + ) + address = relay.start() + unauthenticated_client = ZmqSparseRefitClient( + address, + timeout_s=5.0, + producer_id=2, + ) + client = ZmqSparseRefitClient( + address, + timeout_s=5.0, + producer_id=3, + api_key="secret", + ) + body = b"compressed sparse payload" + checksum = sparse_payload_checksum(body) + try: + with pytest.raises(RuntimeError, match="authentication failed"): + _send_zmq_payload(unauthenticated_client, 6, body) + first = _send_zmq_payload(client, 7, body) + assert first["ok"] + assert first["receiver_total_s"] == 0.25 + assert [len(items) for items in received] == [1, 1] + for items in received: + headers, posted_body = items[0] + assert posted_body == body + assert headers[G_VLLM_REFIT_TRANSFER_HEADER] == "transfer-a" + assert headers[G_VLLM_REFIT_PRODUCER_HEADER] == "3" + assert headers[G_VLLM_REFIT_PAYLOAD_HEADER] == "7" + assert headers[G_VLLM_REFIT_CHECKSUM_HEADER] == checksum + assert headers["x-nemo-rl-refit-key"] == "secret" + + with pytest.raises(RuntimeError, match="checksum mismatch"): + _send_zmq_payload(client, 8, body, "0" * 32) + finally: + unauthenticated_client.close() + client.close() + relay.close() + for server, thread in receivers: + server.shutdown() + thread.join(timeout=5.0) + server.server_close() diff --git a/tests/unit/weight_sync/test_weight_synchronizer.py b/tests/unit/weight_sync/test_weight_synchronizer.py index edec20b6c0a..37def6c9421 100644 --- a/tests/unit/weight_sync/test_weight_synchronizer.py +++ b/tests/unit/weight_sync/test_weight_synchronizer.py @@ -34,6 +34,9 @@ from nemo_rl.weight_sync.ipc_weight_synchronizer import ( IPCWeightSynchronizer, ) +from nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer import ( + VllmRemoteSparseWeightSynchronizer, +) # --------------------------------------------------------------------------- # Helpers @@ -76,6 +79,25 @@ def _mock_cluster(world_size=4, ip="127.0.0.1", port=29500): return cluster +def _remote_sparse_sync( + mock_ray: MagicMock, + transport: str, + stream_result: list[dict[str, int]] | RuntimeError, +) -> tuple[VllmRemoteSparseWeightSynchronizer, MagicMock, MagicMock]: + policy = MagicMock() + policy.init_remote_sparse_delta_baseline.return_value = [MagicMock()] + policy.stream_remote_sparse_weights.return_value = [MagicMock()] + policy.finish_remote_sparse_delta_sync.return_value = [MagicMock()] + generation = MagicMock() + generation.report_refit_server_base_urls.return_value = ["http://receiver"] + generation.start_zmq_sparse_refit_relays.return_value = ["tcp://relay:19090"] + generation.invalidate_kv_cache.return_value = True + mock_ray.get.side_effect = [None, stream_result] + sync = VllmRemoteSparseWeightSynchronizer(policy, generation, transport=transport) + sync.init_communicator() + return sync, policy, generation + + # --------------------------------------------------------------------------- # WeightSynchronizer ABC contract # --------------------------------------------------------------------------- @@ -216,6 +238,109 @@ def test_zero_env_ratio_raises(self, mock_ray, monkeypatch): sync._compute_buffer_size() +class TestVllmRemoteSparseWeightSynchronizer: + @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") + def test_fails_before_transfer_when_kv_cache_invalidation_fails(self, _mock_ray): + policy = MagicMock() + generation = MagicMock() + generation.invalidate_kv_cache.return_value = False + sync = VllmRemoteSparseWeightSynchronizer(policy, generation, transport="s3") + + with pytest.raises(RuntimeError, match="KV cache invalidation failed"): + sync.sync_weights() + policy.stream_remote_sparse_weights.assert_not_called() + + @patch( + "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.flush_vllm_refit_urls" + ) + @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") + def test_initializes_streams_commits_and_updates_baseline( + self, mock_ray, flush, capsys + ): + sync, policy, generation = _remote_sparse_sync( + mock_ray, + "zmq", + [{"payloads": 3, "changed_elements": 3, "total_elements": 100}], + ) + flush.return_value = [ + { + "verification_candidates": 4, + "verification_samples": 4, + "verification_exact_mismatches": 1, + "verification_mismatches": 0, + "verification_abs_sum": 1e-9, + "verification_max_abs": 1e-9, + } + ] + metrics = sync.sync_weights() + + policy.init_remote_sparse_delta_baseline.assert_called_once_with("zmq") + generation.start_zmq_sparse_refit_relays.assert_called_once_with( + ["http://receiver"] + ) + policy.stream_remote_sparse_weights.assert_called_once() + flush.assert_called_once_with( + ["http://receiver"], api_key_env_var=None, timeout_s=600.0 + ) + policy.finish_remote_sparse_delta_sync.assert_called_once_with(True) + assert ( + "REFIT_ZMQ_DELTA_CHANGE changed_elements=3 total_elements=100 " + "changed_pct=3" in capsys.readouterr().out + ) + assert metrics["delta/changed_pct"] == 3.0 + assert metrics["delta_verify/candidates"] == 4.0 + assert metrics["delta_verify/samples"] == 4.0 + assert metrics["delta_verify/exact_mismatches"] == 1.0 + assert metrics["delta_verify/mismatches"] == 0.0 + assert metrics["delta_verify/mean_abs"] == 2.5e-10 + assert metrics["delta_verify/max_abs"] == 1e-9 + assert metrics["transfer/payloads"] == 3.0 + assert not sync.is_stale + + @patch( + "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.flush_vllm_refit_urls" + ) + @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") + def test_sample_mismatch_does_not_commit_baseline(self, mock_ray, flush): + sync, policy, _ = _remote_sparse_sync( + mock_ray, + "zmq", + [{"payloads": 3, "changed_elements": 3, "total_elements": 100}], + ) + flush.return_value = [ + { + "verification_samples": 4, + "verification_mismatches": 1, + "verification_abs_sum": 0.5, + "verification_max_abs": 0.5, + } + ] + + with pytest.raises(RuntimeError, match="1 mismatched deltas out of 4"): + sync.sync_weights() + + policy.finish_remote_sparse_delta_sync.assert_called_once_with(False) + + @patch( + "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.flush_vllm_refit_urls" + ) + @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") + def test_failure_drains_receivers_without_committing_baseline( + self, mock_ray, flush + ): + sync, policy, _ = _remote_sparse_sync( + mock_ray, "s3", RuntimeError("stream failed") + ) + + with pytest.raises(RuntimeError, match="stream failed"): + sync.sync_weights() + + flush.assert_called_once_with( + ["http://receiver"], api_key_env_var=None, timeout_s=60.0 + ) + policy.finish_remote_sparse_delta_sync.assert_called_once_with(False) + + # --------------------------------------------------------------------------- # HTTPWeightSynchronizer # --------------------------------------------------------------------------- From b7160b5ee3dd6df9c7610995104d3c2268b70aef Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Thu, 9 Jul 2026 14:03:13 -0700 Subject: [PATCH 03/18] add more test cases Signed-off-by: Hollow Man --- .../models/generation/test_vllm_generation.py | 224 ++++++++++++- .../test_weight_transfer_remote_sparse.py | 312 +++++++++++++++++- .../weight_sync/test_weight_synchronizer.py | 34 ++ 3 files changed, 568 insertions(+), 2 deletions(-) diff --git a/tests/unit/models/generation/test_vllm_generation.py b/tests/unit/models/generation/test_vllm_generation.py index 5a6e5dc579c..69add935d26 100644 --- a/tests/unit/models/generation/test_vllm_generation.py +++ b/tests/unit/models/generation/test_vllm_generation.py @@ -54,6 +54,18 @@ ) from nemo_rl.models.policy import LoRAConfig, PolicyConfig from nemo_rl.models.policy.lm_policy import Policy +from nemo_rl.utils.weight_transfer_remote_sparse import ( + G_VLLM_REFIT_API_KEY_HEADER, + G_VLLM_REFIT_FLUSH_PATH, + G_VLLM_REFIT_S3_MANIFEST_PATH, +) +from nemo_rl.utils.weight_transfer_zmq import ( + G_VLLM_REFIT_CHECKSUM_HEADER, + G_VLLM_REFIT_PAYLOAD_HEADER, + G_VLLM_REFIT_PRODUCER_HEADER, + G_VLLM_REFIT_TRANSFER_HEADER, + G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, +) model_name = "Qwen/Qwen3-0.6B" # Define basic vLLM test config @@ -275,6 +287,42 @@ def test_sparse_refit_queue_does_not_deduplicate_failed_enqueue() -> None: assert worker._refit_apply_pending_payloads == [] +def test_sparse_refit_collective_response_merges_verification_metrics() -> None: + response = BaseVllmGenerationWorker._refit_collective_response( + [ + { + "receiver_total_s": 1.0, + "verification_candidates": 4, + "verification_samples": 2, + "verification_exact_mismatches": 1, + "verification_mismatches": 0, + "verification_abs_sum": 0.25, + "verification_max_abs": 0.25, + }, + { + "receiver_total_s": 2.0, + "verification_candidates": 4, + "verification_samples": 3, + "verification_exact_mismatches": 2, + "verification_mismatches": 1, + "verification_abs_sum": 0.5, + "verification_max_abs": 0.4, + }, + ] + ) + + assert response == { + "ok": True, + "receiver_total_s": 2.0, + "verification_candidates": 4, + "verification_samples": 5, + "verification_exact_mismatches": 3, + "verification_mismatches": 1, + "verification_abs_sum": 0.75, + "verification_max_abs": 0.4, + } + + def test_sparse_refit_queue_releases_condition_while_backpressured() -> None: first: Future[dict[str, Any]] = Future() second: Future[dict[str, Any]] = Future() @@ -392,6 +440,177 @@ async def collective_rpc( assert response == {"ok": True, "receiver_total_s": 1.0, "payloads": 2} +@pytest.mark.asyncio +async def test_sparse_refit_payload_handlers_decode_and_enqueue(monkeypatch) -> None: + worker = BaseVllmGenerationWorker.__new__(BaseVllmGenerationWorker) + enqueue = MagicMock( + side_effect=[ + {"ok": True, "payloads": 1}, + {"ok": True, "payloads": 1}, + ] + ) + worker._enqueue_sparse_payload_apply = enqueue + monkeypatch.setattr( + "nemo_rl.models.generation.vllm.vllm_worker.download_s3_refit_payload", + lambda _manifest: b"s3-payload", + ) + monkeypatch.setattr( + "nemo_rl.models.generation.vllm.vllm_worker.decode_sparse_payload", + lambda _body, _checksum: b"zmq-payload", + ) + + s3_result = await worker._apply_s3_manifest_payload( + {"key": "object-key", "checksum": "checksum"} + ) + assert s3_result["ok"] + assert s3_result["receiver_s3_download_s"] >= 0.0 + + invalid_request = types.SimpleNamespace(headers={}, body=AsyncMock()) + with pytest.raises(ValueError, match="Missing or invalid"): + await worker._apply_zmq_payload(invalid_request) + + request = types.SimpleNamespace( + headers={ + G_VLLM_REFIT_TRANSFER_HEADER: "transfer", + G_VLLM_REFIT_PRODUCER_HEADER: "2", + G_VLLM_REFIT_PAYLOAD_HEADER: "3", + G_VLLM_REFIT_CHECKSUM_HEADER: "checksum", + }, + body=AsyncMock(return_value=b"compressed"), + ) + zmq_result = await worker._apply_zmq_payload(request) + assert zmq_result["ok"] + assert zmq_result["receiver_zmq_decode_s"] >= 0.0 + assert enqueue.call_args_list == [ + call(b"s3-payload", ("object-key", -1, -1), "checksum"), + call(b"zmq-payload", ("transfer", 2, 3), "checksum"), + ] + + +def test_sparse_refit_api_auth_dispatch_and_error_mapping(monkeypatch) -> None: + from fastapi import FastAPI + from fastapi.testclient import TestClient + + monkeypatch.setenv("NRL_TEST_REFIT_KEY", "secret") + worker = BaseVllmGenerationWorker.__new__(BaseVllmGenerationWorker) + worker.cfg = { + "vllm_cfg": { + "async_engine": True, + "http_refit_api_key_env_var": "NRL_TEST_REFIT_KEY", + } + } + worker._refit_async_loop = None + worker._apply_s3_manifest_payload = AsyncMock( + return_value={"ok": True, "payloads": 1} + ) + worker._apply_zmq_payload = AsyncMock(side_effect=RuntimeError("apply failed")) + worker._flush_queued_sparse_payloads = MagicMock( + return_value={"ok": True, "payloads": 2} + ) + app = FastAPI() + worker._setup_vllm_refit_api_server(app) + headers = {G_VLLM_REFIT_API_KEY_HEADER: "secret"} + + with TestClient(app) as client: + unauthorized = client.post(G_VLLM_REFIT_S3_MANIFEST_PATH, json={}) + s3_response = client.post( + G_VLLM_REFIT_S3_MANIFEST_PATH, + json={"key": "key"}, + headers=headers, + ) + flush_response = client.post(G_VLLM_REFIT_FLUSH_PATH, headers=headers) + zmq_response = client.post( + G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, + content=b"payload", + headers=headers, + ) + + assert unauthorized.status_code == 403 + assert unauthorized.json() == {"ok": False, "error": "unauthorized"} + assert s3_response.status_code == 200 + assert s3_response.json() == {"ok": True, "payloads": 1} + assert flush_response.status_code == 200 + assert flush_response.json() == {"ok": True, "payloads": 2} + assert zmq_response.status_code == 500 + assert zmq_response.json() == {"ok": False, "error": "apply failed"} + assert worker._refit_async_loop is not None + worker._apply_s3_manifest_payload.assert_awaited_once_with({"key": "key"}) + worker._apply_zmq_payload.assert_awaited_once() + worker._flush_queued_sparse_payloads.assert_called_once_with() + + +def test_sync_sparse_refit_server_shutdown_cleans_transport_resources( + monkeypatch, +) -> None: + import uvicorn + + from nemo_rl.models.generation.vllm import vllm_worker as worker_module + + configs = [] + servers = [] + + def make_config(app, **kwargs): + config = types.SimpleNamespace(app=app, **kwargs) + configs.append(config) + return config + + class Server: + def __init__(self, config) -> None: + self.config = config + self.should_exit = False + self.ran = threading.Event() + servers.append(self) + + def run(self) -> None: + self.ran.set() + + monkeypatch.setattr(uvicorn, "Config", make_config) + monkeypatch.setattr(uvicorn, "Server", Server) + monkeypatch.setattr(worker_module, "_get_free_port_local", lambda *_args: 12345) + monkeypatch.setattr(worker_module, "_get_node_ip_local", lambda: "10.0.0.1") + collect = MagicMock() + empty_cache = MagicMock() + monkeypatch.setattr(worker_module.gc, "collect", collect) + monkeypatch.setattr(worker_module.torch.cuda, "empty_cache", empty_cache) + + worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) + worker.cfg = { + "vllm_cfg": {"http_refit_server_port": None}, + "port_range_low": 10000, + "port_range_high": 11000, + } + worker._setup_vllm_refit_api_server = MagicMock() + worker._refit_http_server = None + worker._setup_vllm_refit_server() + + assert len(configs) == 1 + assert configs[0].host == "0.0.0.0" + assert configs[0].port == 12345 + assert servers[0].ran.wait(timeout=1.0) + assert worker.report_refit_server_base_url() == "http://10.0.0.1:12345" + worker._setup_vllm_refit_api_server.assert_called_once_with(configs[0].app) + + relay = MagicMock() + llm = MagicMock() + worker._zmq_refit_server = (relay, "tcp://relay") + worker._flush_queued_sparse_payloads = MagicMock() + worker._refit_apply_executor = MagicMock() + worker.llm = llm + worker.tokenizer = object() + + assert worker.shutdown() is True + relay.close.assert_called_once_with() + worker._flush_queued_sparse_payloads.assert_called_once_with() + worker._refit_apply_executor.shutdown.assert_called_once_with(wait=True) + assert servers[0].should_exit is True + assert worker._refit_http_server is None + llm.collective_rpc.assert_called_once_with("cleanup", args=()) + assert worker.llm is None + assert worker.tokenizer is None + collect.assert_called_once_with() + empty_cache.assert_called_once_with() + + @pytest.mark.asyncio async def test_async_sparse_refit_post_init_records_worker_locality() -> None: worker = VllmAsyncGenerationWorkerImpl.__new__(VllmAsyncGenerationWorkerImpl) @@ -405,7 +624,10 @@ async def test_async_sparse_refit_post_init_records_worker_locality() -> None: assert worker.vllm_device_ids == ["0"] assert worker._refit_workers_share_node is True - worker.llm.collective_rpc.assert_awaited_once_with("report_node_hostname", args=()) + assert worker.llm.collective_rpc.await_args_list == [ + call("bind_numa", args=()), + call("report_node_hostname", args=()), + ] def test_async_sparse_refit_exposes_zmq_relay(monkeypatch) -> None: diff --git a/tests/unit/utils/test_weight_transfer_remote_sparse.py b/tests/unit/utils/test_weight_transfer_remote_sparse.py index 906ae092d4f..cf48be816ec 100644 --- a/tests/unit/utils/test_weight_transfer_remote_sparse.py +++ b/tests/unit/utils/test_weight_transfer_remote_sparse.py @@ -21,7 +21,7 @@ import torch import zstandard -from nemo_rl.utils import weight_transfer_remote_sparse +from nemo_rl.utils import weight_transfer_remote_sparse, weight_transfer_zmq from nemo_rl.utils.weight_transfer_remote_sparse import download_s3_refit_payload from nemo_rl.utils.weight_transfer_sparse_codec import ( DeltaCompressionTracker, @@ -216,6 +216,316 @@ def fail_transfer(_body, _payload_index): assert exported == list(range(4)) +def test_sparse_baseline_snapshots_only_owned_export_chunks( + monkeypatch, capsys +) -> None: + monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "4") + + class Tracker: + sparse_bucket_size_bytes = 4 + + def __init__(self) -> None: + self.names = [] + + def snapshot_baseline(self, chunk) -> None: + self.names.extend(name for name, _tensor in chunk) + + tracker = Tracker() + weight_transfer_remote_sparse.init_sparse_delta_baseline_from_iterator( + [(f"weight-{index}", torch.tensor([float(index)])) for index in range(4)], + delta_tracker=tracker, + shard_rank=1, + shard_count=2, + transport="zmq", + ) + + assert tracker.names == ["weight-1", "weight-3"] + assert "chunks=4" in capsys.readouterr().out + + +def test_s3_manifest_transport_validates_configuration(monkeypatch) -> None: + kwargs = { + "iterator": (), + "delta_tracker": SimpleNamespace(), + "transfer_id": "transfer", + "api_key_env_var": None, + "timeout_s": 1.0, + "shard_rank": 0, + "shard_count": 1, + } + monkeypatch.setenv("NRL_REFIT_S3_BUCKET", "bucket") + with pytest.raises(ValueError, match="URL is required"): + weight_transfer_remote_sparse.stream_sparse_delta_payloads_via_s3_manifest( + refit_targets=[], **kwargs + ) + + monkeypatch.delenv("NRL_REFIT_S3_BUCKET") + with pytest.raises(RuntimeError, match="NRL_REFIT_S3_BUCKET"): + weight_transfer_remote_sparse.stream_sparse_delta_payloads_via_s3_manifest( + refit_targets=["http://receiver"], **kwargs + ) + + +def test_s3_manifest_transport_uploads_notifies_and_deletes(monkeypatch) -> None: + operations = [] + posts = [] + + class Store: + bucket = "bucket" + region = "us-west-2" + + def put(self, key, body) -> None: + operations.append(("put", key, body)) + + def delete(self, key) -> None: + operations.append(("delete", key)) + + store = Store() + monkeypatch.setenv("NRL_REFIT_S3_BUCKET", store.bucket) + monkeypatch.setenv("NRL_REFIT_S3_REGION", store.region) + monkeypatch.setenv("NRL_REFIT_S3_PREFIX", "/prefix/") + monkeypatch.setenv("NRL_REFIT_S3_UPLOAD_WORKERS", "3") + monkeypatch.setenv("NRL_TEST_REFIT_KEY", "secret") + monkeypatch.setattr( + weight_transfer_remote_sparse, + "_get_manifest_s3_store", + lambda *_args: store, + ) + + def post(endpoints, manifest, **kwargs): + posts.append((endpoints, manifest, kwargs)) + return [ + {"ok": True, "receiver_total_s": 1.0}, + {"ok": True, "receiver_total_s": 2.0}, + ] + + def stream(iterator, **kwargs): + assert list(iterator) == [("weight", torch.tensor([1.0]))] + assert kwargs["transport"] == "s3" + assert kwargs["transfer_workers"] == 3 + response = kwargs["send_payload"](b"payload", 3) + assert response["receiver"] == {"receiver_total_s": 2.0} + return {"payloads": 1, "changed_elements": 1, "total_elements": 1} + + monkeypatch.setattr( + weight_transfer_remote_sparse, "post_vllm_refit_endpoints", post + ) + monkeypatch.setattr( + weight_transfer_remote_sparse, "stream_sparse_delta_payloads", stream + ) + + result = weight_transfer_remote_sparse.stream_sparse_delta_payloads_via_s3_manifest( + [("weight", torch.tensor([1.0]))], + delta_tracker=SimpleNamespace(), + refit_targets=[" http://receiver-a/ ", "http://receiver-b"], + transfer_id="transfer", + api_key_env_var="NRL_TEST_REFIT_KEY", + timeout_s=7.0, + shard_rank=2, + shard_count=4, + ) + + key = "prefix/transfer/000002/000003.pt" + assert result == {"payloads": 1, "changed_elements": 1, "total_elements": 1} + assert operations == [("put", key, b"payload"), ("delete", key)] + assert posts == [ + ( + [ + "http://receiver-a/nemo-rl/refit/s3-manifest", + "http://receiver-b/nemo-rl/refit/s3-manifest", + ], + { + "bucket": store.bucket, + "region": store.region, + "key": key, + "checksum": sparse_payload_checksum(b"payload"), + }, + {"api_key": "secret", "timeout_s": 7.0}, + ) + ] + + +def test_zmq_stream_routes_shards_and_reuses_clients(monkeypatch) -> None: + created = [] + sent = [] + monkeypatch.setenv("NRL_TEST_REFIT_KEY", "secret") + monkeypatch.setenv("NRL_REFIT_ZMQ_SEND_WORKERS", "2") + monkeypatch.delattr(weight_transfer_zmq._ZMQ_LOCAL, "clients", raising=False) + + class Client: + def __init__(self, address, **kwargs) -> None: + created.append((address, kwargs)) + + def send_payload(self, **kwargs): + sent.append(kwargs) + return {"ok": True, "receiver_total_s": 0.5} + + def stream(_iterator, **kwargs): + assert kwargs["transport"] == "zmq" + assert kwargs["transfer_workers"] == 2 + for payload_id in range(2): + response = kwargs["send_payload"](f"body-{payload_id}".encode(), payload_id) + assert response["receiver"]["ok"] + return {"payloads": 2, "changed_elements": 2, "total_elements": 2} + + monkeypatch.setattr(weight_transfer_zmq, "ZmqSparseRefitClient", Client) + monkeypatch.setattr(weight_transfer_zmq, "stream_sparse_delta_payloads", stream) + + with pytest.raises(ValueError, match="address is required"): + weight_transfer_zmq.stream_sparse_delta_payloads_via_zmq( + (), + delta_tracker=SimpleNamespace(), + refit_targets=[], + transfer_id="transfer", + api_key_env_var=None, + timeout_s=1.0, + shard_rank=0, + shard_count=1, + ) + + kwargs = { + "iterator": (), + "delta_tracker": SimpleNamespace(), + "refit_targets": ["tcp://receiver-a", " tcp://receiver-b "], + "transfer_id": "transfer", + "api_key_env_var": "NRL_TEST_REFIT_KEY", + "timeout_s": 7.0, + "shard_rank": 3, + "shard_count": 4, + } + assert ( + weight_transfer_zmq.stream_sparse_delta_payloads_via_zmq(**kwargs)["payloads"] + == 2 + ) + assert ( + weight_transfer_zmq.stream_sparse_delta_payloads_via_zmq(**kwargs)["payloads"] + == 2 + ) + + assert created == [ + ( + "tcp://receiver-b", + {"timeout_s": 7.0, "producer_id": 3, "api_key": "secret"}, + ) + ] + assert [item["payload_id"] for item in sent] == [0, 1, 0, 1] + assert all( + item["checksum"] == sparse_payload_checksum(item["body"]) for item in sent + ) + + +def test_zmq_client_filters_replies_rejects_nack_and_retries(monkeypatch) -> None: + class Socket: + def __init__(self, replies=(), send_failures=0) -> None: + self.replies = list(replies) + self.send_failures = send_failures + self.sent = [] + + def send_multipart(self, frames, **_kwargs) -> None: + self.sent.append(frames) + if self.send_failures: + self.send_failures -= 1 + raise weight_transfer_zmq.zmq.Again() + + def poll(self, *_args) -> bool: + return bool(self.replies) + + def recv_multipart(self): + return self.replies.pop(0) + + def client(socket) -> ZmqSparseRefitClient: + result = ZmqSparseRefitClient.__new__(ZmqSparseRefitClient) + result._address = "tcp://receiver" + result._timeout_ms = 1000 + result._producer_id = 4 + result._api_key = "secret" + result._socket = socket + return result + + success = { + "ok": True, + "transfer_id": "transfer-a", + "producer_id": 4, + "payload_id": 7, + } + socket = Socket( + [ + [b"malformed"], + [ + b"ACK", + json.dumps({**success, "transfer_id": "stale"}).encode(), + ], + [b"ACK", json.dumps(success).encode()], + ] + ) + assert _send_zmq_payload(client(socket), 7, b"body") == success + assert json.loads(socket.sent[0][1])["api_key"] == "secret" + + denied = Socket([[b"NACK", json.dumps({"ok": False, "error": "denied"}).encode()]]) + with pytest.raises(RuntimeError, match="denied"): + _send_zmq_payload(client(denied), 7, b"body") + + monkeypatch.setenv("NRL_REFIT_ZMQ_RETRIES", "1") + with pytest.raises(TimeoutError, match="payload 7"): + _send_zmq_payload(client(Socket(send_failures=2)), 7, b"body") + + +def test_zmq_server_rejects_malformed_messages() -> None: + server = ZmqSparseRefitServer.__new__(ZmqSparseRefitServer) + server._token = "secret" + body = b"body" + metadata = { + "protocol": "nemo-rl-sparse-zmq-v1", + "api_key": "secret", + "transfer_id": "transfer", + "producer_id": 0, + "payload_id": 1, + "checksum": sparse_payload_checksum(body), + } + + with pytest.raises(ValueError, match="Expected 4"): + server._parse_data_message([]) + with pytest.raises(ValueError, match="Unsupported ZeroMQ sparse refit message"): + server._parse_data_message( + [b"id", b"OTHER", json.dumps(metadata).encode(), body] + ) + with pytest.raises(ValueError, match="protocol"): + server._parse_data_message( + [ + b"id", + b"DATA", + json.dumps({**metadata, "protocol": "other"}).encode(), + body, + ] + ) + with pytest.raises(PermissionError, match="authentication"): + server._parse_data_message( + [ + b"id", + b"DATA", + json.dumps({**metadata, "api_key": "wrong"}).encode(), + body, + ] + ) + with pytest.raises(ValueError, match="identity"): + server._parse_data_message( + [b"id", b"DATA", json.dumps({**metadata, "transfer_id": ""}).encode(), body] + ) + with pytest.raises(ValueError, match="checksum mismatch"): + server._parse_data_message( + [ + b"id", + b"DATA", + json.dumps({**metadata, "checksum": "wrong"}).encode(), + body, + ] + ) + + assert server._parse_data_message( + [b"id", b"DATA", json.dumps(metadata).encode(), body] + )[1] == ("transfer", 0, 1) + + def _receiver_server(received): class Handler(BaseHTTPRequestHandler): def do_POST(self): diff --git a/tests/unit/weight_sync/test_weight_synchronizer.py b/tests/unit/weight_sync/test_weight_synchronizer.py index 37def6c9421..99b593ec9d9 100644 --- a/tests/unit/weight_sync/test_weight_synchronizer.py +++ b/tests/unit/weight_sync/test_weight_synchronizer.py @@ -239,6 +239,40 @@ def test_zero_env_ratio_raises(self, mock_ray, monkeypatch): class TestVllmRemoteSparseWeightSynchronizer: + def test_init_communicator_requires_receiver_endpoints(self): + policy = MagicMock() + generation = MagicMock() + generation.report_refit_server_base_urls.return_value = [] + sync = VllmRemoteSparseWeightSynchronizer(policy, generation, transport="s3") + + with pytest.raises(ValueError, match="endpoints are missing"): + sync.init_communicator() + + @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") + def test_shutdown_cancels_pending_work_and_stops_zmq(self, mock_ray): + policy = MagicMock() + generation = MagicMock() + sync = VllmRemoteSparseWeightSynchronizer(policy, generation, transport="zmq") + init_ref, commit_ref = MagicMock(), MagicMock() + sync._baseline_init_refs = [init_ref] + sync._baseline_commit_refs = [commit_ref] + sync._refit_urls = ["http://receiver"] + sync._targets = ["tcp://relay"] + sync._stale = False + + sync.mark_stale() + sync.shutdown() + + assert mock_ray.cancel.call_count == 2 + mock_ray.cancel.assert_any_call(init_ref, force=False) + mock_ray.cancel.assert_any_call(commit_ref, force=False) + generation.stop_zmq_sparse_refit_relays.assert_called_once_with() + assert sync.is_stale + assert sync._baseline_init_refs is None + assert sync._baseline_commit_refs is None + assert sync._refit_urls == [] + assert sync._targets == [] + @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") def test_fails_before_transfer_when_kv_cache_invalidation_fails(self, _mock_ray): policy = MagicMock() From bfd4259c218f520d684ed6b71f342cda3825d346 Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Thu, 9 Jul 2026 20:37:01 -0700 Subject: [PATCH 04/18] minimizing the code surface for easier maintainance Signed-off-by: Hollow Man --- docs/design-docs/sparse-delta-refit.md | 6 + nemo_rl/algorithms/grpo.py | 44 +- nemo_rl/distributed/virtual_cluster.py | 3 +- .../models/generation/vllm/vllm_backend.py | 789 ++---------------- .../models/generation/vllm/vllm_generation.py | 43 +- .../generation/vllm/vllm_sparse_delta.py | 761 +++++++++++++++++ .../generation/vllm/vllm_sparse_refit.py | 445 ++++++++++ nemo_rl/models/generation/vllm/vllm_worker.py | 416 +-------- .../generation/vllm/vllm_worker_async.py | 24 +- nemo_rl/models/policy/lm_policy.py | 49 +- .../policy/workers/megatron_policy_worker.py | 80 +- .../workers/megatron_remote_sparse_refit.py | 87 ++ .../utils/weight_transfer_remote_sparse.py | 4 + .../vllm_remote_sparse_weight_synchronizer.py | 83 +- pyrefly.toml | 3 + tests/unit/algorithms/test_grpo.py | 16 +- .../models/generation/test_vllm_backend.py | 291 +------ .../models/generation/test_vllm_generation.py | 502 +---------- .../generation/test_vllm_sparse_delta.py | 288 +++++++ .../generation/test_vllm_sparse_refit.py | 491 +++++++++++ .../test_megatron_remote_sparse_refit.py | 60 ++ .../models/policy/test_megatron_worker.py | 62 -- ..._vllm_remote_sparse_weight_synchronizer.py | 215 +++++ .../weight_sync/test_weight_synchronizer.py | 159 ---- 24 files changed, 2574 insertions(+), 2347 deletions(-) create mode 100644 nemo_rl/models/generation/vllm/vllm_sparse_delta.py create mode 100644 nemo_rl/models/generation/vllm/vllm_sparse_refit.py create mode 100644 nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py create mode 100644 tests/unit/models/generation/test_vllm_sparse_delta.py create mode 100644 tests/unit/models/generation/test_vllm_sparse_refit.py create mode 100644 tests/unit/models/policy/test_megatron_remote_sparse_refit.py create mode 100644 tests/unit/weight_sync/test_vllm_remote_sparse_weight_synchronizer.py diff --git a/docs/design-docs/sparse-delta-refit.md b/docs/design-docs/sparse-delta-refit.md index 8e46292316d..743a9fc31ea 100644 --- a/docs/design-docs/sparse-delta-refit.md +++ b/docs/design-docs/sparse-delta-refit.md @@ -6,6 +6,12 @@ encoding, compression, backpressure, receiver apply, and transactional baseline commit logic. Payload checksums and transfer-scoped IDs make HTTP retries idempotent. Policy workers commit their baselines only after every receiver flush succeeds. + +The implementation is opt-in. Core policy and vLLM workers contain only lazy +delegates; export/encoding, transport, receiver queuing, and sparse placement +live in dedicated modules and are not initialized by existing IPC, HTTP, or +NCCL refit paths. + On a fresh run, generation starts from the shared checkpoint while policy workers build the CPU baseline asynchronously; the first transfer follows the first optimizer step. Resumed runs synchronize before generation. diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index afa5938d390..0599d2e2662 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -118,9 +118,6 @@ from nemo_rl.utils.nsys import maybe_gpu_profile_step from nemo_rl.utils.timer import TimeoutChecker, Timer from nemo_rl.utils.venvs import create_local_venv_on_each_node -from nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer import ( - VllmRemoteSparseWeightSynchronizer, -) # =============================================================================== # Configuration @@ -910,9 +907,7 @@ def _spinup_nemo_gym(base_urls, model_name): # vllm model loading prefers clean environment, initialize policy_generation before policy in colocated mode backend = generation_config["backend"] generation_config["model_name"] = policy_config["model_name"] # Needed for vLLM - refit_transport = ( - generation_config.get("refit_transport") if backend == "vllm" else None - ) + refit_transport = None # Dictionary to store worker initialization timing stats for logging worker_init_timing_metrics = {} @@ -1101,23 +1096,18 @@ def initialize_generation_with_policy( elif backend == "vllm": # vLLM generation: setup config, then initialize with policy generation_config = cast(VllmConfig, generation_config) - vllm_cfg = generation_config["vllm_cfg"] - if refit_transport not in (None, "vllm_s3_sparse", "vllm_zmq_sparse"): - raise ValueError(f"Unsupported vLLM refit transport {refit_transport!r}.") - if refit_transport is not None: - if ( - colocated_inference - or not policy_config["megatron_cfg"]["enabled"] - or vllm_cfg["precision"] == "fp8" - or vllm_cfg["kv_cache_dtype"].startswith("fp8") - or not generation_config.get("delta_compression") - or generation_config.get("quant_cfg") - or generation_config.get("real_quant") - ): - raise ValueError( - f"{refit_transport} requires a non-colocated Megatron policy, " - "BF16/FP16 vLLM, delta compression, and an unquantized rollout." - ) + if generation_config.get("refit_transport") is not None: + # Keep optional remote transport dependencies off the default path. + from nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer import ( + VllmRemoteSparseWeightSynchronizer, + validate_vllm_remote_sparse_refit, + ) + + refit_transport = validate_vllm_remote_sparse_refit( + generation_config, + colocated=colocated_inference, + megatron_enabled=policy_config["megatron_cfg"]["enabled"], + ) if generation_config["vllm_cfg"]["precision"] == "fp8": assert loss_config.use_importance_sampling_correction, ( @@ -2087,12 +2077,10 @@ def refit_policy_generation( Returns: Scalar metrics reported by the selected weight synchronizer. """ - if ( - isinstance(policy_generation, VllmGeneration) - and policy_generation.weight_synchronizer is not None - ): + synchronizer = getattr(policy_generation, "weight_synchronizer", None) + if synchronizer is not None: return ( - policy_generation.weight_synchronizer.sync_weights( + synchronizer.sync_weights( timer=timer, kv_scales=kv_scales, ) diff --git a/nemo_rl/distributed/virtual_cluster.py b/nemo_rl/distributed/virtual_cluster.py index 52926991745..c4e80445070 100644 --- a/nemo_rl/distributed/virtual_cluster.py +++ b/nemo_rl/distributed/virtual_cluster.py @@ -1095,5 +1095,4 @@ def __del__(self) -> None: the cluster is lost due to leaving a function scope. It's always recommended that the user calls shutdown(). """ - if not sys.is_finalizing(): - self.shutdown() + self.shutdown() diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index 0437767a309..5cadc893252 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -12,11 +12,9 @@ # See the License for the specific language governing permissions and # limitations under the License. import gc -import io import re -import time +import socket import traceback -from dataclasses import dataclass from typing import Any, cast import torch @@ -27,15 +25,9 @@ calculate_aligned_size, rebuild_cuda_tensor_from_ipc, ) -from nemo_rl.utils import weight_transfer_sparse_codec as sparse_codec from nemo_rl.utils.nsys import wrap_with_nvtx_name from nemo_rl.utils.packed_tensor import packed_broadcast_consumer -_EXPERT_WEIGHT_RE = re.compile( - r"^(?P.*\.experts)\.(?P\d+)\." - r"(?Pgate_proj|up_proj|down_proj)\.weight$" -) - try: import vllm # noqa: F401 except ImportError: @@ -100,31 +92,8 @@ def _read_mtp_layer_weights_from_checkpoint( return weights -@dataclass(frozen=True) -class _SparseDeltaTargetPlan: - target: torch.Tensor | None - source_shape: tuple[int, ...] = () - source_strides: tuple[int, ...] = () - target_strides: tuple[int, ...] = () - target_offset: int = 0 - shard_dim: int | None = None - shard_start: int = 0 - shard_size: int = 0 - segment_shards: tuple[tuple[int, int, int], ...] = () - log_delta_transform: bool = False - identity: bool = False - - class VllmInternalWorkerExtension: - state_dict_info: dict[str, Any] | None = None - _direct_sparse_delta_targets: dict[str, torch.Tensor] | None = None - _direct_sparse_delta_plan_cache: dict[str, _SparseDeltaTargetPlan | None] | None = ( - None - ) - _direct_sparse_delta_verification: ( - list[tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]] | None - ) = None - _direct_sparse_delta_verification_candidates = 0 + _sparse_delta_applier: Any = None def bind_numa(self) -> bool: """Pin this TP worker to its GPU's NUMA-local CPUs/memory. @@ -173,8 +142,6 @@ def report_device_id(self) -> str: def report_node_hostname(self) -> str: """Return the host shared by worker processes on this node.""" - import socket - return socket.gethostname() def get_zmq_address(self): @@ -205,30 +172,9 @@ def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: e.g. {tensor_name: (shape, dtype)} """ 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 - self._direct_sparse_delta_targets = None - self._direct_sparse_delta_plan_cache = None - self._direct_sparse_delta_verification = [] - self._direct_sparse_delta_verification_candidates = 0 - - def _process_weights_after_loading( - self, - model_config: Any, - target_device: torch.device, - ) -> None: - from vllm.config import set_current_vllm_config - from vllm.model_executor.model_loader.utils import ( - process_weights_after_loading, - ) - - with set_current_vllm_config(self.model_runner.vllm_config): - process_weights_after_loading( - self.model_runner.model, - model_config, - target_device, - ) def _maybe_process_fp8_kv_cache(self) -> None: - """Process weights after loading for FP8 KV cache static scales.""" + """Process weights after loading for FP8 KV cache (static scales).""" use_fp8_kv_cache = False if hasattr(self.model_runner.vllm_config, "cache_config"): kv_cache_dtype = getattr( @@ -237,13 +183,27 @@ def _maybe_process_fp8_kv_cache(self) -> None: use_fp8_kv_cache = ( kv_cache_dtype is not None and "fp8" in str(kv_cache_dtype).lower() ) + if not use_fp8_kv_cache: return - self._process_weights_after_loading( - self.model_runner.model_config, - next(self.model_runner.model.parameters()).device, + + # 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]], @@ -402,576 +362,19 @@ def _load_weights(self, weights): self._load_draft_weights(draft_weights) - def _apply_sparse_weight_deltas( - self, - payload_tensors: tuple[torch.Tensor, torch.Tensor], - metadata: list[dict[str, Any]], - ) -> None: - """Apply sparse deltas directly after validating every target plan.""" - architectures = self.model_runner.vllm_config.model_config.architectures - from nemo_rl.models.generation.vllm.quantization import fp8 - - if {"GptOssForCausalLM", "Gemma3ForConditionalGeneration"} & set( - architectures - ) or fp8.is_fp8_model(self.model_runner.vllm_config): - raise RuntimeError( - "Direct sparse delta refit does not support transformed or FP8 weights." - ) - - if self._direct_sparse_delta_targets is None: - model = self.model_runner.model - self._direct_sparse_delta_targets = dict(model.named_parameters()) | dict( - model.named_buffers() + def _get_sparse_delta_applier(self) -> Any: + if self._sparse_delta_applier is None: + # Avoid importing sparse placement code for existing refit transports. + from nemo_rl.models.generation.vllm.vllm_sparse_delta import ( + VllmSparseDeltaApplier, ) - targets = self._direct_sparse_delta_targets - raw_locations, raw_values = payload_tensors - plan_cache = self._direct_sparse_delta_plan_cache - if plan_cache is None: - plan_cache = self._direct_sparse_delta_plan_cache = {} - plans = [] - for item in metadata: - name = str(item["name"]) - if name not in plan_cache: - plan_cache[name] = self._direct_sparse_delta_target_plan(item, targets) - plan = plan_cache[name] - if plan is None: - raise RuntimeError( - f"No direct sparse delta target plan for {item['name']!r}." - ) - plans.append((item, plan)) - - with torch.no_grad(): - for item, plan in plans: - target = plan.target - verification_locations = item.get("verification_locations", []) - self._direct_sparse_delta_verification_candidates += len( - verification_locations - ) - if target is None: - continue - - if verification_locations and not plan.log_delta_transform: - sample_locations, sample_deltas = ( - self._local_sparse_delta_update_inputs( - torch.tensor(verification_locations, device=target.device), - torch.tensor( - item["verification_deltas"], - device=target.device, - dtype=target.dtype, - ), - plan, - ) - ) - if sample_locations.numel(): - before = target.data.view(-1).index_select(0, sample_locations) - expected_delta = ( - before + sample_deltas - ).float() - before.float() - verification = self._direct_sparse_delta_verification - if verification is None: - verification = self._direct_sparse_delta_verification = [] - verification.append( - ( - target, - sample_locations, - before.float(), - expected_delta, - ) - ) - value_start = int(item["value_start"]) - value_end = int(item["value_end"]) - values = raw_values[value_start:value_end].to( - device=target.device, - dtype=target.dtype, - non_blocking=True, - ) - if plan.identity and item["index_encoding"] == "range": - range_start = int(item["range_start"]) - range_count = value_end - value_start - target.data.view(-1).narrow(0, range_start, range_count).add_( - values - ) - else: - locations = sparse_codec.sparse_locations_for_item( - item, - raw_locations, - device=target.device, - ) - locations, values = self._local_sparse_delta_update_inputs( - locations, - values, - plan, - ) - if locations.numel(): - target_flat = target.data.view(-1) - if plan.log_delta_transform: - current = target_flat.index_select(0, locations) - updated = current * values.float().exp().to( - dtype=current.dtype - ) - target_flat.index_copy_(0, locations, updated) - else: - target_flat.index_add_(0, locations, values) - - def _direct_sparse_delta_module( - self, - target: torch.Tensor, - module_name: str, - ) -> Any: - loader = getattr(target, "weight_loader", None) - return getattr( - loader, "__self__", None - ) or self.model_runner.model.get_submodule(module_name) - - def _direct_sparse_delta_target_plan( - self, - item: dict[str, Any], - targets: dict[str, torch.Tensor], - ) -> _SparseDeltaTargetPlan | None: - name = str(item["name"]) - if name.startswith("mtp."): - return _SparseDeltaTargetPlan(target=None) - mapper = getattr(self.model_runner.model, "hf_to_vllm_mapper", None) - target_name = cast(Any, mapper)._map_name(name) if mapper is not None else name - if target_name is None or target_name.startswith("draft."): - return None - if ".mixer." in target_name: - mamba_plan = self._direct_sparse_delta_mamba2_plan( - item, target_name, targets - ) - if mamba_plan is not None: - return mamba_plan - if ".mixer.conv1d." in target_name or ".mixer.in_proj." in target_name: - return None - if any(f".{candidate}_proj." in target_name for candidate in ("q", "k", "v")): - return self._direct_sparse_delta_qkv_plan(item, target_name, targets) - if _EXPERT_WEIGHT_RE.match(target_name): - return self._direct_sparse_delta_expert_plan(item, target_name, targets) - if any(f".{candidate}_proj." in target_name for candidate in ("gate", "up")): - merged_plan = self._direct_sparse_delta_merged_column_plan( - item, target_name, targets + self._sparse_delta_applier = VllmSparseDeltaApplier( + self.model_runner, + self.device, + rank=int(getattr(self, "rank", 0)), ) - if merged_plan is not None: - return merged_plan - - target = targets.get(target_name) - if target is None: - return None - source_shape = tuple(item["shape"]) - target_shape = tuple(target.shape) - if target_shape == source_shape: - return self._make_sparse_delta_target_plan(target, source_shape) - return self._direct_sparse_delta_shard_plan(item, target) - - def _direct_sparse_delta_qkv_plan( - self, - item: dict[str, Any], - target_name: str, - targets: dict[str, torch.Tensor], - ) -> _SparseDeltaTargetPlan | None: - shard_id = next(x for x in "qkv" if f".{x}_proj." in target_name) - packed_name = target_name.replace(f".{shard_id}_proj.", ".qkv_proj.", 1) - target = targets.get(packed_name) - if target is None: - return None - output_dim = int(cast(Any, target).output_dim) % target.ndim - module = self._direct_sparse_delta_module(target, packed_name.rsplit(".", 1)[0]) - shard_offset = int(module._get_shard_offset_mapping(shard_id)) - shard_size = int(module._get_shard_size_mapping(shard_id)) - shard_rank = int(module.tp_rank) - if shard_id != "q": - shard_rank //= int(module.num_kv_head_replicas) - - source_shape = tuple(item["shape"]) - shard_start = shard_rank * shard_size - if source_shape[output_dim] < shard_start: - return _SparseDeltaTargetPlan(target=None) - return self._make_sparse_delta_target_plan( - target, - source_shape=source_shape, - shard_dim=output_dim, - shard_start=shard_start, - shard_size=min(shard_size, source_shape[output_dim] - shard_start), - target_offset=shard_offset * target.stride(output_dim), - ) - - def _direct_sparse_delta_merged_column_plan( - self, - item: dict[str, Any], - target_name: str, - targets: dict[str, torch.Tensor], - ) -> _SparseDeltaTargetPlan | None: - projection = next( - candidate - for candidate in ("gate", "up") - if f".{candidate}_proj." in target_name - ) - shard_id = 0 if projection == "gate" else 1 - packed_name = target_name.replace(f".{projection}_proj.", ".gate_up_proj.", 1) - target = targets.get(packed_name) - output_dim = getattr(target, "output_dim", None) - if target is None or not isinstance(output_dim, int): - return None - - output_dim %= target.ndim - module = self._direct_sparse_delta_module(target, packed_name.rsplit(".", 1)[0]) - output_sizes = tuple(int(size) for size in module.output_sizes) - tp_size = int(module.tp_size) - source_shape = tuple(item["shape"]) - if ( - shard_id >= len(output_sizes) - or tp_size < 1 - or output_sizes[shard_id] % tp_size - or output_dim >= len(source_shape) - or source_shape[output_dim] != output_sizes[shard_id] - ): - return None - - shard_size = output_sizes[shard_id] // tp_size - target_start = sum(output_sizes[:shard_id]) // tp_size - if target.shape[output_dim] < target_start + shard_size: - return None - shard_start = int(module.tp_rank) * shard_size - return self._make_sparse_delta_target_plan( - target, - source_shape=source_shape, - shard_dim=output_dim, - shard_start=shard_start, - shard_size=shard_size, - target_offset=target_start * target.stride(output_dim), - ) - - def _direct_sparse_delta_mamba2_plan( - self, - item: dict[str, Any], - target_name: str, - targets: dict[str, torch.Tensor], - ) -> _SparseDeltaTargetPlan | None: - target = targets.get(target_name) - if target is None: - return None - - if target_name.endswith(".A"): - source_shape = tuple(item["shape"]) - if tuple(target.shape) == source_shape: - return self._make_sparse_delta_target_plan( - target, - source_shape=source_shape, - log_delta_transform=True, - ) - return self._direct_sparse_delta_shard_plan( - item, - target, - log_delta_transform=True, - ) - - if not (".mixer.conv1d." in target_name or ".mixer.in_proj." in target_name): - return None - - mixer_name = target_name.split(".mixer.", 1)[0] + ".mixer" - attrs = cast(Any, self.model_runner.model.get_submodule(mixer_name)) - tp_size = int(attrs.tp_size) - if tp_size <= 1: - return None - intermediate_size = int(attrs.intermediate_size) - groups_ssm_state_size = int(attrs.groups_ssm_state_size) - num_heads = int(attrs.num_heads) - source_shape = tuple(item["shape"]) - fixed_size = intermediate_size - if ".mixer.in_proj." in target_name: - fixed_size += intermediate_size + num_heads - group_size, remainder = divmod(source_shape[0] - fixed_size, 2) - extra_group_size = groups_ssm_state_size - group_size - if remainder or group_size <= 0 or extra_group_size < 0: - return None - tp_rank = int( - getattr(target, "tp_rank", int(getattr(self, "rank", 0)) % tp_size) - ) - intermediate = (intermediate_size, 0, False) - group = (groups_ssm_state_size, extra_group_size, extra_group_size > 0) - segment_specs = ( - (intermediate, group, group) - if ".mixer.conv1d." in target_name - else (intermediate, intermediate, group, group, (num_heads, 0, False)) - ) - - target_shape = tuple(target.shape) - source_to_target_dims = tuple(range(len(source_shape))) - if len(target_shape) == len(source_shape) + 1 and target_shape[1] == 1: - source_to_target_dims = (0, *range(2, len(target_shape))) - elif len(target_shape) != len(source_shape): - return None - segment_shards: list[tuple[int, int, int]] = [] - target_start = 0 - source_start = 0 - for full_dim, extra, duplicate_groups in segment_specs: - shard_size = full_dim // tp_size - rank = 0 if duplicate_groups else tp_rank - source_dim = full_dim - extra - source_local_start = source_start + rank * shard_size - take = min(shard_size, source_dim - rank * shard_size) - if take > 0: - segment_shards.append((source_local_start, target_start, take)) - target_start += shard_size - source_start += source_dim - if source_shape[0] != source_start or target_shape[0] != target_start: - return None - - return self._make_sparse_delta_target_plan( - target, - source_shape=source_shape, - source_to_target_dims=source_to_target_dims, - shard_dim=0, - segment_shards=tuple(segment_shards), - ) - - def _direct_sparse_delta_expert_plan( - self, - item: dict[str, Any], - target_name: str, - targets: dict[str, torch.Tensor], - ) -> _SparseDeltaTargetPlan | None: - match = cast(re.Match[str], _EXPERT_WEIGHT_RE.match(target_name)) - - prefix = match.group("prefix") - global_expert_id = int(match.group("expert")) - proj = match.group("proj") - packed_weight, shard_id = { - "gate_proj": ("w13_weight", "w1"), - "up_proj": ("w13_weight", "w3"), - "down_proj": ("w2_weight", "w2"), - }[proj] - packed_name = f"{prefix}.{packed_weight}" - - target = targets.get(packed_name) - if target is None: - return None - module_attrs = self._direct_sparse_delta_module( - target, packed_name.rsplit(".", 1)[0] - ) - if shard_id == "w3" and not module_attrs.moe_config.is_act_and_mul: - shard_id = "w1" - local_expert_id = int( - module_attrs._map_global_expert_id_to_local_expert_id(global_expert_id) - ) - if local_expert_id < 0: - return _SparseDeltaTargetPlan(target=None) - - source_shape = tuple(item["shape"]) - target_shape = tuple(target.shape) - shard_dim = 1 if shard_id == "w2" else 0 - if local_expert_id >= target_shape[0]: - return None - target_shard_dim = shard_dim + 1 - - shard_size = target_shape[target_shard_dim] - if shard_id in ("w1", "w3") and module_attrs.moe_config.is_act_and_mul: - shard_size //= 2 - target_shard_offset = shard_size if shard_id == "w3" else 0 - if target_shape[target_shard_dim] < target_shard_offset + shard_size: - return None - tp_rank = int(module_attrs.tp_rank) - shard_start = tp_rank * shard_size - if source_shape[shard_dim] < shard_start: - return _SparseDeltaTargetPlan(target=None) - - return self._make_sparse_delta_target_plan( - target, - source_shape=source_shape, - source_to_target_dims=tuple(dim + 1 for dim in range(len(source_shape))), - target_offset=( - local_expert_id * target.stride(0) - + target_shard_offset * target.stride(target_shard_dim) - ), - shard_dim=shard_dim, - shard_start=shard_start, - shard_size=min(shard_size, source_shape[shard_dim] - shard_start), - ) - - def _make_sparse_delta_target_plan( - self, - target: torch.Tensor, - source_shape: tuple[int, ...], - *, - source_to_target_dims: tuple[int, ...] | None = None, - target_offset: int = 0, - shard_dim: int | None = None, - shard_start: int = 0, - shard_size: int = 0, - segment_shards: tuple[tuple[int, int, int], ...] = (), - log_delta_transform: bool = False, - ) -> _SparseDeltaTargetPlan | None: - if source_to_target_dims is None: - source_to_target_dims = tuple(range(len(source_shape))) - target_shape = tuple(target.shape) - ignored_dim = ( - shard_dim if shard_dim is not None else 0 if segment_shards else -1 - ) - if len(source_to_target_dims) != len(source_shape) or any( - target_dim >= len(target_shape) - or ( - source_dim != ignored_dim - and source_shape[source_dim] != target_shape[target_dim] - ) - for source_dim, target_dim in enumerate(source_to_target_dims) - ): - return None - identity = ( - shard_dim is None - and target_offset == 0 - and not segment_shards - and not log_delta_transform - and source_to_target_dims == tuple(range(len(source_shape))) - and source_shape == target_shape - ) - return _SparseDeltaTargetPlan( - target=target, - source_shape=source_shape, - source_strides=torch.empty(source_shape, device="meta").stride(), - target_strides=tuple( - target.stride(target_dim) for target_dim in source_to_target_dims - ), - target_offset=target_offset, - shard_dim=shard_dim, - shard_start=shard_start, - shard_size=shard_size, - segment_shards=segment_shards, - log_delta_transform=log_delta_transform, - identity=identity, - ) - - def _direct_sparse_delta_shard_plan( - self, - item: dict[str, Any], - target: torch.Tensor, - *, - log_delta_transform: bool = False, - ) -> _SparseDeltaTargetPlan | None: - source_shape = tuple(item["shape"]) - target_shape = tuple(target.shape) - if len(source_shape) != len(target_shape): - return None - - candidate_dims = list( - dict.fromkeys( - dim % len(source_shape) - for attr in ("output_dim", "input_dim") - if isinstance(dim := getattr(target, attr, None), int) - ) - ) - if not candidate_dims: - candidate_dims = [ - dim - for dim, (source_dim, target_dim) in enumerate( - zip(source_shape, target_shape, strict=True) - ) - if source_dim != target_dim - ] - if len(candidate_dims) != 1: - return None - - for shard_dim in candidate_dims: - shard_size = target_shape[shard_dim] - tp_size = int(getattr(target, "tp_size", 1)) - if tp_size <= 1: - if shard_size <= 0 or source_shape[shard_dim] % shard_size: - continue - tp_size = source_shape[shard_dim] // shard_size - if tp_size <= 1: - continue - if source_shape[shard_dim] > shard_size * tp_size: - continue - tp_rank = int( - getattr(target, "tp_rank", int(getattr(self, "rank", 0)) % tp_size) - ) - plan = self._make_sparse_delta_target_plan( - target=target, - source_shape=source_shape, - shard_dim=shard_dim, - shard_start=tp_rank * shard_size, - shard_size=shard_size, - log_delta_transform=log_delta_transform, - ) - if plan is not None: - return plan - return None - - def _local_sparse_delta_update_inputs( - self, - locations: torch.Tensor, - values: torch.Tensor, - plan: _SparseDeltaTargetPlan, - ) -> tuple[torch.Tensor, torch.Tensor]: - if plan.identity: - return locations, values - - source_shape = plan.source_shape - source_strides = plan.source_strides - target_strides = plan.target_strides - shard_dim = plan.shard_dim - - if source_strides == target_strides: - if shard_dim is None: - return locations + plan.target_offset, values - if shard_dim == 0: - shard_stride = source_strides[0] - shard_coords = torch.div(locations, shard_stride, rounding_mode="floor") - if plan.segment_shards: - mapped_locations = locations + plan.target_offset - keep = torch.zeros_like(locations, dtype=torch.bool) - for source_start, target_start, take in plan.segment_shards: - segment = (shard_coords >= source_start) & ( - shard_coords < source_start + take - ) - mapped_locations[segment] += ( - target_start - source_start - ) * shard_stride - keep |= segment - return mapped_locations[keep], values[keep] - shard_end = min( - plan.shard_start + plan.shard_size, source_shape[shard_dim] - ) - keep = (shard_coords >= plan.shard_start) & (shard_coords < shard_end) - return ( - locations[keep] - + plan.target_offset - - plan.shard_start * shard_stride, - values[keep], - ) - - selected_locations = locations - selected_values = values - - if shard_dim is not None: - shard_coords = torch.div( - locations, - source_strides[shard_dim], - rounding_mode="floor", - ).remainder(source_shape[shard_dim]) - shard_end = min( - plan.shard_start + plan.shard_size, - source_shape[shard_dim], - ) - keep = (shard_coords >= plan.shard_start) & (shard_coords < shard_end) - selected_locations = locations[keep] - selected_values = values[keep] - if selected_locations.numel() == 0: - return selected_locations, selected_values - - local_locations = torch.full_like(selected_locations, plan.target_offset) - for dim, (source_stride, target_stride) in enumerate( - zip(source_strides, target_strides, strict=True) - ): - coord = torch.div( - selected_locations, - source_stride, - rounding_mode="floor", - ).remainder(source_shape[dim]) - if dim == plan.shard_dim: - coord = coord - plan.shard_start - local_locations.add_(coord * target_stride) - return local_locations, selected_values + return self._sparse_delta_applier @wrap_with_nvtx_name("vllm_internal_worker_extension/update_weights_via_ipc_zmq") def update_weights_via_ipc_zmq(self) -> bool: @@ -991,24 +394,26 @@ def update_weights_via_ipc_zmq(self) -> bool: if payload == IPCProtocol.COMPLETE: # means the update is done - self._process_weights_after_loading(self.model_config, self.device) + from vllm.config import set_current_vllm_config + from vllm.model_executor.model_loader.utils import ( + process_weights_after_loading, + ) + + with set_current_vllm_config(self.model_runner.vllm_config): + process_weights_after_loading( + self.model_runner.model, self.model_config, self.device + ) self.zmq_socket.send(IPCProtocol.ACK.value.encode()) break ipc_handle, list_keys, used_bytes = payload buffer = rebuild_cuda_tensor_from_ipc(ipc_handle, self.device.index) - state_dict_info = self.state_dict_info - if state_dict_info is None: - raise RuntimeError( - "state_dict_info is not prepared. " - "Call prepare_refit_info before loading weights." - ) weight = None weights = [] offset = 0 for key in list_keys: - shape, dtype = state_dict_info[key] + shape, dtype = self.state_dict_info[key] # pyrefly if isinstance(shape, list): shape = torch.Size(shape) @@ -1100,125 +505,29 @@ def update_weights_from_collective(self) -> bool: torch.cuda.empty_cache() return True - @wrap_with_nvtx_name( - "vllm_internal_worker_extension/update_weights_from_serialized_sparse_payload" - ) def update_weights_from_serialized_sparse_payload( self, serialized_payload: bytes, ) -> dict[str, Any]: - """Apply one serialized sparse-delta payload.""" - return self._load_and_apply_sparse_payload(io.BytesIO(serialized_payload)) - - def _load_and_apply_sparse_payload( - self, - source: str | io.BytesIO, - ) -> dict[str, Any]: - started = time.perf_counter() - payload = cast( - sparse_codec.TensorPayload, - torch.load( - source, - map_location="cpu", - weights_only=True, - ), + return self._get_sparse_delta_applier().update_weights_from_serialized_sparse_payload( + serialized_payload ) - deserialize_s = time.perf_counter() - started - result = self._apply_sparse_request(payload) - result["receiver_deserialize_s"] = deserialize_s - result["receiver_total_s"] = time.perf_counter() - started - return result - @wrap_with_nvtx_name( - "vllm_internal_worker_extension/update_weights_from_sparse_payload_files" - ) def update_weights_from_sparse_payload_files( self, *payload_paths: str, ) -> dict[str, Any]: - """Apply sparse payloads in FIFO order.""" - started = time.perf_counter() - deserialize_s = 0.0 - sparse_apply_s = 0.0 - for path in payload_paths: - result = self._load_and_apply_sparse_payload(path) - deserialize_s += float(result["receiver_deserialize_s"]) - sparse_apply_s += float(result["receiver_sparse_apply_s"]) - return { - "ok": True, - "receiver_deserialize_s": deserialize_s, - "receiver_sparse_apply_s": sparse_apply_s, - "receiver_total_s": time.perf_counter() - started, - } - - def _apply_sparse_request( - self, - payload: sparse_codec.TensorPayload, - ) -> dict[str, Any]: - locations, values, metadata = payload - - sparse_started = time.perf_counter() - self._apply_sparse_weight_deltas((locations, values), metadata) - sparse_apply_s = time.perf_counter() - sparse_started - - return { - "ok": True, - "receiver_sparse_apply_s": sparse_apply_s, - } + return ( + self._get_sparse_delta_applier().update_weights_from_sparse_payload_files( + *payload_paths + ) + ) def synchronize_device(self) -> None: - """Synchronize this vLLM worker's CUDA device after deferred refit applies.""" - if torch.cuda.is_available(): - torch.cuda.synchronize(self.device) + self._get_sparse_delta_applier().synchronize_device() def finish_sparse_delta_refit(self) -> dict[str, Any]: - """Synchronize and compare bounded producer samples with applied weights.""" - self.synchronize_device() - verification = self._direct_sparse_delta_verification or [] - candidates = self._direct_sparse_delta_verification_candidates - self._direct_sparse_delta_verification = [] - self._direct_sparse_delta_verification_candidates = 0 - if not verification: - return { - "ok": True, - "verification_candidates": candidates, - "verification_samples": 0, - "verification_exact_mismatches": 0, - "verification_mismatches": 0, - "verification_abs_sum": 0.0, - "verification_max_abs": 0.0, - } - - with torch.no_grad(): - actual_delta = torch.cat( - [ - target.data.view(-1).index_select(0, locations).float() - before - for target, locations, before, _ in verification - ] - ) - expected_delta = torch.cat([expected for _, _, _, expected in verification]) - difference = (actual_delta - expected_delta).abs() - exact_mismatches = actual_delta.ne(expected_delta) - mismatches = ~torch.isclose( - actual_delta, expected_delta, rtol=1e-6, atol=1e-8 - ) - stats = torch.stack( - ( - difference.sum(), - difference.max(), - exact_mismatches.sum().float(), - mismatches.sum().float(), - ) - ).cpu() - return { - "ok": True, - "verification_candidates": candidates, - "verification_samples": actual_delta.numel(), - "verification_exact_mismatches": int(stats[2]), - "verification_mismatches": int(stats[3]), - "verification_abs_sum": float(stats[0]), - "verification_max_abs": float(stats[1]), - } + return self._get_sparse_delta_applier().finish_sparse_delta_refit() def cleanup(self) -> None: """Shutdown and cleanup resources.""" diff --git a/nemo_rl/models/generation/vllm/vllm_generation.py b/nemo_rl/models/generation/vllm/vllm_generation.py index e51942d1259..e889fe9ed2d 100644 --- a/nemo_rl/models/generation/vllm/vllm_generation.py +++ b/nemo_rl/models/generation/vllm/vllm_generation.py @@ -15,7 +15,6 @@ import asyncio import logging import os -import sys import warnings from collections import defaultdict from typing import ( @@ -916,45 +915,22 @@ def shutdown(self) -> bool: def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: """Prepare the info for refit.""" + # Choose the appropriate method based on async_engine setting method_name = ( "prepare_refit_info_async" if self.cfg["vllm_cfg"]["async_engine"] else "prepare_refit_info" ) - self._run_refit_workers(method_name, state_dict_info=state_dict_info) - - def report_refit_server_base_urls(self) -> list[str]: - """Return base URLs for vLLM workers exposing sparse refit endpoints.""" - return [ - url - for url in self._run_refit_workers("report_refit_server_base_url") - if url - ] - def _run_refit_workers(self, method_name: str, **kwargs: Any) -> list[Any]: - if not self.worker_group or not self.worker_group.workers: - raise RuntimeError("Worker group is not initialized") - return ray.get( - self.worker_group.run_all_workers_single_data( - method_name, - run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], - **kwargs, - ) + # Use run_all_workers_single_data to send data to all workers + futures = self.worker_group.run_all_workers_single_data( + method_name, + state_dict_info=state_dict_info, + run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], ) - def start_zmq_sparse_refit_relays(self, refit_urls: list[str]) -> list[str]: - """Start one ZeroMQ relay per vLLM replica and return TCP addresses.""" - return [ - address - for address in self._run_refit_workers( - "start_zmq_sparse_refit_relay", refit_urls=refit_urls - ) - if address - ] - - def stop_zmq_sparse_refit_relays(self) -> None: - if self.worker_group and self.worker_group.workers: - self._run_refit_workers("stop_zmq_sparse_refit_relay") + # Wait for all futures to complete + ray.get(futures) def update_weights_via_ipc_zmq(self) -> list[ray.ObjectRef]: """Update weights of the policy using IPC handles via ZMQ socket.""" @@ -1080,8 +1056,7 @@ def __del__(self) -> None: the object is lost due to leaving a function scope. It's always recommended that the user calls shutdown(). """ - if not sys.is_finalizing(): - self.shutdown() + self.shutdown() def invalidate_kv_cache(self) -> bool: """Invalidate reusable caches in vLLM (e.g., prefix/KV cache) after weight updates. diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py new file mode 100644 index 00000000000..d26497aa1d6 --- /dev/null +++ b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py @@ -0,0 +1,761 @@ +# 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. + +"""Direct sparse-delta placement and application for vLLM workers.""" + +import io +import re +import time +from dataclasses import dataclass +from typing import Any, cast + +import torch + +from nemo_rl.utils import weight_transfer_sparse_codec as sparse_codec +from nemo_rl.utils.nsys import wrap_with_nvtx_name + +_EXPERT_WEIGHT_RE = re.compile( + r"^(?P.*\.experts)\.(?P\d+)\." + r"(?Pgate_proj|up_proj|down_proj)\.weight$" +) + + +@dataclass(frozen=True) +class _SparseDeltaTargetPlan: + target: torch.Tensor | None + source_shape: tuple[int, ...] = () + source_strides: tuple[int, ...] = () + target_strides: tuple[int, ...] = () + target_offset: int = 0 + shard_dim: int | None = None + shard_start: int = 0 + shard_size: int = 0 + segment_shards: tuple[tuple[int, int, int], ...] = () + log_delta_transform: bool = False + identity: bool = False + + +class VllmSparseDeltaApplier: + """Own sparse placement state without extending the normal refit path.""" + + def __init__( + self, + model_runner: Any, + device: torch.device, + *, + rank: int = 0, + ) -> None: + self.model_runner = model_runner + self._cuda_device_index = device.index + self.rank = rank + self._direct_sparse_delta_targets: dict[str, torch.Tensor] | None = None + self._direct_sparse_delta_plan_cache: dict[ + str, _SparseDeltaTargetPlan | None + ] = {} + self._direct_sparse_delta_verification: list[ + tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] + ] = [] + self._direct_sparse_delta_verification_candidates = 0 + + def _apply_sparse_weight_deltas( + self, + payload_tensors: tuple[torch.Tensor, torch.Tensor], + metadata: list[dict[str, Any]], + ) -> None: + """Apply sparse deltas directly after validating every target plan.""" + architectures = self.model_runner.vllm_config.model_config.architectures + # Delay the vLLM-dependent FP8 helper until a payload is applied. + from nemo_rl.models.generation.vllm.quantization import fp8 + + if {"GptOssForCausalLM", "Gemma3ForConditionalGeneration"} & set( + architectures + ) or fp8.is_fp8_model(self.model_runner.vllm_config): + raise RuntimeError( + "Direct sparse delta refit does not support transformed or FP8 weights." + ) + + if self._direct_sparse_delta_targets is None: + model = self.model_runner.model + self._direct_sparse_delta_targets = dict(model.named_parameters()) | dict( + model.named_buffers() + ) + targets = self._direct_sparse_delta_targets + raw_locations, raw_values = payload_tensors + plan_cache = self._direct_sparse_delta_plan_cache + if plan_cache is None: + plan_cache = self._direct_sparse_delta_plan_cache = {} + plans = [] + for item in metadata: + name = str(item["name"]) + if name not in plan_cache: + plan_cache[name] = self._direct_sparse_delta_target_plan(item, targets) + plan = plan_cache[name] + if plan is None: + raise RuntimeError( + f"No direct sparse delta target plan for {item['name']!r}." + ) + plans.append((item, plan)) + + with torch.no_grad(): + for item, plan in plans: + target = plan.target + verification_locations = item.get("verification_locations", []) + self._direct_sparse_delta_verification_candidates += len( + verification_locations + ) + if target is None: + continue + + if verification_locations and not plan.log_delta_transform: + sample_locations, sample_deltas = ( + self._local_sparse_delta_update_inputs( + torch.tensor(verification_locations, device=target.device), + torch.tensor( + item["verification_deltas"], + device=target.device, + dtype=target.dtype, + ), + plan, + ) + ) + if sample_locations.numel(): + before = target.data.view(-1).index_select(0, sample_locations) + expected_delta = ( + before + sample_deltas + ).float() - before.float() + verification = self._direct_sparse_delta_verification + if verification is None: + verification = self._direct_sparse_delta_verification = [] + verification.append( + ( + target, + sample_locations, + before.float(), + expected_delta, + ) + ) + + value_start = int(item["value_start"]) + value_end = int(item["value_end"]) + values = raw_values[value_start:value_end].to( + device=target.device, + dtype=target.dtype, + non_blocking=True, + ) + if plan.identity and item["index_encoding"] == "range": + range_start = int(item["range_start"]) + range_count = value_end - value_start + target.data.view(-1).narrow(0, range_start, range_count).add_( + values + ) + else: + locations = sparse_codec.sparse_locations_for_item( + item, + raw_locations, + device=target.device, + ) + locations, values = self._local_sparse_delta_update_inputs( + locations, + values, + plan, + ) + if locations.numel(): + target_flat = target.data.view(-1) + if plan.log_delta_transform: + current = target_flat.index_select(0, locations) + updated = current * values.float().exp().to( + dtype=current.dtype + ) + target_flat.index_copy_(0, locations, updated) + else: + target_flat.index_add_(0, locations, values) + + def _direct_sparse_delta_module( + self, + target: torch.Tensor, + module_name: str, + ) -> Any: + loader = getattr(target, "weight_loader", None) + return getattr( + loader, "__self__", None + ) or self.model_runner.model.get_submodule(module_name) + + def _direct_sparse_delta_target_plan( + self, + item: dict[str, Any], + targets: dict[str, torch.Tensor], + ) -> _SparseDeltaTargetPlan | None: + name = str(item["name"]) + if name.startswith("mtp."): + return _SparseDeltaTargetPlan(target=None) + mapper = getattr(self.model_runner.model, "hf_to_vllm_mapper", None) + target_name = cast(Any, mapper)._map_name(name) if mapper is not None else name + if target_name is None or target_name.startswith("draft."): + return None + if ".mixer." in target_name: + mamba_plan = self._direct_sparse_delta_mamba2_plan( + item, target_name, targets + ) + if mamba_plan is not None: + return mamba_plan + if ".mixer.conv1d." in target_name or ".mixer.in_proj." in target_name: + return None + if any(f".{candidate}_proj." in target_name for candidate in ("q", "k", "v")): + return self._direct_sparse_delta_qkv_plan(item, target_name, targets) + if _EXPERT_WEIGHT_RE.match(target_name): + return self._direct_sparse_delta_expert_plan(item, target_name, targets) + if any(f".{candidate}_proj." in target_name for candidate in ("gate", "up")): + merged_plan = self._direct_sparse_delta_merged_column_plan( + item, target_name, targets + ) + if merged_plan is not None: + return merged_plan + + target = targets.get(target_name) + if target is None: + return None + source_shape = tuple(item["shape"]) + target_shape = tuple(target.shape) + if target_shape == source_shape: + return self._make_sparse_delta_target_plan(target, source_shape) + return self._direct_sparse_delta_shard_plan(item, target) + + def _direct_sparse_delta_qkv_plan( + self, + item: dict[str, Any], + target_name: str, + targets: dict[str, torch.Tensor], + ) -> _SparseDeltaTargetPlan | None: + shard_id = next(x for x in "qkv" if f".{x}_proj." in target_name) + packed_name = target_name.replace(f".{shard_id}_proj.", ".qkv_proj.", 1) + target = targets.get(packed_name) + if target is None: + return None + output_dim = int(cast(Any, target).output_dim) % target.ndim + module = self._direct_sparse_delta_module(target, packed_name.rsplit(".", 1)[0]) + shard_offset = int(module._get_shard_offset_mapping(shard_id)) + shard_size = int(module._get_shard_size_mapping(shard_id)) + shard_rank = int(module.tp_rank) + if shard_id != "q": + shard_rank //= int(module.num_kv_head_replicas) + + source_shape = tuple(item["shape"]) + shard_start = shard_rank * shard_size + if source_shape[output_dim] < shard_start: + return _SparseDeltaTargetPlan(target=None) + return self._make_sparse_delta_target_plan( + target, + source_shape=source_shape, + shard_dim=output_dim, + shard_start=shard_start, + shard_size=min(shard_size, source_shape[output_dim] - shard_start), + target_offset=shard_offset * target.stride(output_dim), + ) + + def _direct_sparse_delta_merged_column_plan( + self, + item: dict[str, Any], + target_name: str, + targets: dict[str, torch.Tensor], + ) -> _SparseDeltaTargetPlan | None: + projection = next( + candidate + for candidate in ("gate", "up") + if f".{candidate}_proj." in target_name + ) + shard_id = 0 if projection == "gate" else 1 + packed_name = target_name.replace(f".{projection}_proj.", ".gate_up_proj.", 1) + target = targets.get(packed_name) + output_dim = getattr(target, "output_dim", None) + if target is None or not isinstance(output_dim, int): + return None + + output_dim %= target.ndim + module = self._direct_sparse_delta_module(target, packed_name.rsplit(".", 1)[0]) + output_sizes = tuple(int(size) for size in module.output_sizes) + tp_size = int(module.tp_size) + source_shape = tuple(item["shape"]) + if ( + shard_id >= len(output_sizes) + or tp_size < 1 + or output_sizes[shard_id] % tp_size + or output_dim >= len(source_shape) + or source_shape[output_dim] != output_sizes[shard_id] + ): + return None + + shard_size = output_sizes[shard_id] // tp_size + target_start = sum(output_sizes[:shard_id]) // tp_size + if target.shape[output_dim] < target_start + shard_size: + return None + shard_start = int(module.tp_rank) * shard_size + return self._make_sparse_delta_target_plan( + target, + source_shape=source_shape, + shard_dim=output_dim, + shard_start=shard_start, + shard_size=shard_size, + target_offset=target_start * target.stride(output_dim), + ) + + def _direct_sparse_delta_mamba2_plan( + self, + item: dict[str, Any], + target_name: str, + targets: dict[str, torch.Tensor], + ) -> _SparseDeltaTargetPlan | None: + target = targets.get(target_name) + if target is None: + return None + + if target_name.endswith(".A"): + source_shape = tuple(item["shape"]) + if tuple(target.shape) == source_shape: + return self._make_sparse_delta_target_plan( + target, + source_shape=source_shape, + log_delta_transform=True, + ) + return self._direct_sparse_delta_shard_plan( + item, + target, + log_delta_transform=True, + ) + + if not (".mixer.conv1d." in target_name or ".mixer.in_proj." in target_name): + return None + + mixer_name = target_name.split(".mixer.", 1)[0] + ".mixer" + attrs = cast(Any, self.model_runner.model.get_submodule(mixer_name)) + tp_size = int(attrs.tp_size) + if tp_size <= 1: + return None + intermediate_size = int(attrs.intermediate_size) + groups_ssm_state_size = int(attrs.groups_ssm_state_size) + num_heads = int(attrs.num_heads) + source_shape = tuple(item["shape"]) + fixed_size = intermediate_size + if ".mixer.in_proj." in target_name: + fixed_size += intermediate_size + num_heads + group_size, remainder = divmod(source_shape[0] - fixed_size, 2) + extra_group_size = groups_ssm_state_size - group_size + if remainder or group_size <= 0 or extra_group_size < 0: + return None + tp_rank = int( + getattr(target, "tp_rank", int(getattr(self, "rank", 0)) % tp_size) + ) + intermediate = (intermediate_size, 0, False) + group = (groups_ssm_state_size, extra_group_size, extra_group_size > 0) + segment_specs = ( + (intermediate, group, group) + if ".mixer.conv1d." in target_name + else (intermediate, intermediate, group, group, (num_heads, 0, False)) + ) + + target_shape = tuple(target.shape) + source_to_target_dims = tuple(range(len(source_shape))) + if len(target_shape) == len(source_shape) + 1 and target_shape[1] == 1: + source_to_target_dims = (0, *range(2, len(target_shape))) + elif len(target_shape) != len(source_shape): + return None + segment_shards: list[tuple[int, int, int]] = [] + target_start = 0 + source_start = 0 + for full_dim, extra, duplicate_groups in segment_specs: + shard_size = full_dim // tp_size + rank = 0 if duplicate_groups else tp_rank + source_dim = full_dim - extra + source_local_start = source_start + rank * shard_size + take = min(shard_size, source_dim - rank * shard_size) + if take > 0: + segment_shards.append((source_local_start, target_start, take)) + target_start += shard_size + source_start += source_dim + if source_shape[0] != source_start or target_shape[0] != target_start: + return None + + return self._make_sparse_delta_target_plan( + target, + source_shape=source_shape, + source_to_target_dims=source_to_target_dims, + shard_dim=0, + segment_shards=tuple(segment_shards), + ) + + def _direct_sparse_delta_expert_plan( + self, + item: dict[str, Any], + target_name: str, + targets: dict[str, torch.Tensor], + ) -> _SparseDeltaTargetPlan | None: + match = cast(re.Match[str], _EXPERT_WEIGHT_RE.match(target_name)) + + prefix = match.group("prefix") + global_expert_id = int(match.group("expert")) + proj = match.group("proj") + packed_weight, shard_id = { + "gate_proj": ("w13_weight", "w1"), + "up_proj": ("w13_weight", "w3"), + "down_proj": ("w2_weight", "w2"), + }[proj] + packed_name = f"{prefix}.{packed_weight}" + + target = targets.get(packed_name) + if target is None: + return None + module_attrs = self._direct_sparse_delta_module( + target, packed_name.rsplit(".", 1)[0] + ) + if shard_id == "w3" and not module_attrs.moe_config.is_act_and_mul: + shard_id = "w1" + local_expert_id = int( + module_attrs._map_global_expert_id_to_local_expert_id(global_expert_id) + ) + if local_expert_id < 0: + return _SparseDeltaTargetPlan(target=None) + + source_shape = tuple(item["shape"]) + target_shape = tuple(target.shape) + shard_dim = 1 if shard_id == "w2" else 0 + if local_expert_id >= target_shape[0]: + return None + target_shard_dim = shard_dim + 1 + + shard_size = target_shape[target_shard_dim] + if shard_id in ("w1", "w3") and module_attrs.moe_config.is_act_and_mul: + shard_size //= 2 + target_shard_offset = shard_size if shard_id == "w3" else 0 + if target_shape[target_shard_dim] < target_shard_offset + shard_size: + return None + tp_rank = int(module_attrs.tp_rank) + shard_start = tp_rank * shard_size + if source_shape[shard_dim] < shard_start: + return _SparseDeltaTargetPlan(target=None) + + return self._make_sparse_delta_target_plan( + target, + source_shape=source_shape, + source_to_target_dims=tuple(dim + 1 for dim in range(len(source_shape))), + target_offset=( + local_expert_id * target.stride(0) + + target_shard_offset * target.stride(target_shard_dim) + ), + shard_dim=shard_dim, + shard_start=shard_start, + shard_size=min(shard_size, source_shape[shard_dim] - shard_start), + ) + + def _make_sparse_delta_target_plan( + self, + target: torch.Tensor, + source_shape: tuple[int, ...], + *, + source_to_target_dims: tuple[int, ...] | None = None, + target_offset: int = 0, + shard_dim: int | None = None, + shard_start: int = 0, + shard_size: int = 0, + segment_shards: tuple[tuple[int, int, int], ...] = (), + log_delta_transform: bool = False, + ) -> _SparseDeltaTargetPlan | None: + if source_to_target_dims is None: + source_to_target_dims = tuple(range(len(source_shape))) + target_shape = tuple(target.shape) + ignored_dim = ( + shard_dim if shard_dim is not None else 0 if segment_shards else -1 + ) + if len(source_to_target_dims) != len(source_shape) or any( + target_dim >= len(target_shape) + or ( + source_dim != ignored_dim + and source_shape[source_dim] != target_shape[target_dim] + ) + for source_dim, target_dim in enumerate(source_to_target_dims) + ): + return None + identity = ( + shard_dim is None + and target_offset == 0 + and not segment_shards + and not log_delta_transform + and source_to_target_dims == tuple(range(len(source_shape))) + and source_shape == target_shape + ) + return _SparseDeltaTargetPlan( + target=target, + source_shape=source_shape, + source_strides=torch.empty(source_shape, device="meta").stride(), + target_strides=tuple( + target.stride(target_dim) for target_dim in source_to_target_dims + ), + target_offset=target_offset, + shard_dim=shard_dim, + shard_start=shard_start, + shard_size=shard_size, + segment_shards=segment_shards, + log_delta_transform=log_delta_transform, + identity=identity, + ) + + def _direct_sparse_delta_shard_plan( + self, + item: dict[str, Any], + target: torch.Tensor, + *, + log_delta_transform: bool = False, + ) -> _SparseDeltaTargetPlan | None: + source_shape = tuple(item["shape"]) + target_shape = tuple(target.shape) + if len(source_shape) != len(target_shape): + return None + + candidate_dims = list( + dict.fromkeys( + dim % len(source_shape) + for attr in ("output_dim", "input_dim") + if isinstance(dim := getattr(target, attr, None), int) + ) + ) + if not candidate_dims: + candidate_dims = [ + dim + for dim, (source_dim, target_dim) in enumerate( + zip(source_shape, target_shape, strict=True) + ) + if source_dim != target_dim + ] + if len(candidate_dims) != 1: + return None + + for shard_dim in candidate_dims: + shard_size = target_shape[shard_dim] + tp_size = int(getattr(target, "tp_size", 1)) + if tp_size <= 1: + if shard_size <= 0 or source_shape[shard_dim] % shard_size: + continue + tp_size = source_shape[shard_dim] // shard_size + if tp_size <= 1: + continue + if source_shape[shard_dim] > shard_size * tp_size: + continue + tp_rank = int( + getattr(target, "tp_rank", int(getattr(self, "rank", 0)) % tp_size) + ) + plan = self._make_sparse_delta_target_plan( + target=target, + source_shape=source_shape, + shard_dim=shard_dim, + shard_start=tp_rank * shard_size, + shard_size=shard_size, + log_delta_transform=log_delta_transform, + ) + if plan is not None: + return plan + return None + + def _local_sparse_delta_update_inputs( + self, + locations: torch.Tensor, + values: torch.Tensor, + plan: _SparseDeltaTargetPlan, + ) -> tuple[torch.Tensor, torch.Tensor]: + if plan.identity: + return locations, values + + source_shape = plan.source_shape + source_strides = plan.source_strides + target_strides = plan.target_strides + shard_dim = plan.shard_dim + + if source_strides == target_strides: + if shard_dim is None: + return locations + plan.target_offset, values + if shard_dim == 0: + shard_stride = source_strides[0] + shard_coords = torch.div(locations, shard_stride, rounding_mode="floor") + if plan.segment_shards: + mapped_locations = locations + plan.target_offset + keep = torch.zeros_like(locations, dtype=torch.bool) + for source_start, target_start, take in plan.segment_shards: + segment = (shard_coords >= source_start) & ( + shard_coords < source_start + take + ) + mapped_locations[segment] += ( + target_start - source_start + ) * shard_stride + keep |= segment + return mapped_locations[keep], values[keep] + shard_end = min( + plan.shard_start + plan.shard_size, source_shape[shard_dim] + ) + keep = (shard_coords >= plan.shard_start) & (shard_coords < shard_end) + return ( + locations[keep] + + plan.target_offset + - plan.shard_start * shard_stride, + values[keep], + ) + + selected_locations = locations + selected_values = values + + if shard_dim is not None: + shard_coords = torch.div( + locations, + source_strides[shard_dim], + rounding_mode="floor", + ).remainder(source_shape[shard_dim]) + shard_end = min( + plan.shard_start + plan.shard_size, + source_shape[shard_dim], + ) + keep = (shard_coords >= plan.shard_start) & (shard_coords < shard_end) + selected_locations = locations[keep] + selected_values = values[keep] + if selected_locations.numel() == 0: + return selected_locations, selected_values + + local_locations = torch.full_like(selected_locations, plan.target_offset) + for dim, (source_stride, target_stride) in enumerate( + zip(source_strides, target_strides, strict=True) + ): + coord = torch.div( + selected_locations, + source_stride, + rounding_mode="floor", + ).remainder(source_shape[dim]) + if dim == plan.shard_dim: + coord = coord - plan.shard_start + local_locations.add_(coord * target_stride) + return local_locations, selected_values + + @wrap_with_nvtx_name( + "vllm_internal_worker_extension/update_weights_from_serialized_sparse_payload" + ) + def update_weights_from_serialized_sparse_payload( + self, + serialized_payload: bytes, + ) -> dict[str, Any]: + """Apply one serialized sparse-delta payload.""" + return self._load_and_apply_sparse_payload(io.BytesIO(serialized_payload)) + + def _load_and_apply_sparse_payload( + self, + source: str | io.BytesIO, + ) -> dict[str, Any]: + started = time.perf_counter() + payload = cast( + sparse_codec.TensorPayload, + torch.load( + source, + map_location="cpu", + weights_only=True, + ), + ) + deserialize_s = time.perf_counter() - started + result = self._apply_sparse_request(payload) + result["receiver_deserialize_s"] = deserialize_s + result["receiver_total_s"] = time.perf_counter() - started + return result + + @wrap_with_nvtx_name( + "vllm_internal_worker_extension/update_weights_from_sparse_payload_files" + ) + def update_weights_from_sparse_payload_files( + self, + *payload_paths: str, + ) -> dict[str, Any]: + """Apply sparse payloads in FIFO order.""" + started = time.perf_counter() + deserialize_s = 0.0 + sparse_apply_s = 0.0 + for path in payload_paths: + result = self._load_and_apply_sparse_payload(path) + deserialize_s += float(result["receiver_deserialize_s"]) + sparse_apply_s += float(result["receiver_sparse_apply_s"]) + return { + "ok": True, + "receiver_deserialize_s": deserialize_s, + "receiver_sparse_apply_s": sparse_apply_s, + "receiver_total_s": time.perf_counter() - started, + } + + def _apply_sparse_request( + self, + payload: sparse_codec.TensorPayload, + ) -> dict[str, Any]: + locations, values, metadata = payload + + sparse_started = time.perf_counter() + self._apply_sparse_weight_deltas((locations, values), metadata) + sparse_apply_s = time.perf_counter() - sparse_started + + return { + "ok": True, + "receiver_sparse_apply_s": sparse_apply_s, + } + + def synchronize_device(self) -> None: + """Synchronize this vLLM worker's CUDA device after deferred refit applies.""" + if torch.cuda.is_available(): + torch.cuda.synchronize(self._cuda_device_index) + + def finish_sparse_delta_refit(self) -> dict[str, Any]: + """Synchronize and compare bounded producer samples with applied weights.""" + self.synchronize_device() + verification = self._direct_sparse_delta_verification or [] + candidates = self._direct_sparse_delta_verification_candidates + self._direct_sparse_delta_verification = [] + self._direct_sparse_delta_verification_candidates = 0 + if not verification: + return { + "ok": True, + "verification_candidates": candidates, + "verification_samples": 0, + "verification_exact_mismatches": 0, + "verification_mismatches": 0, + "verification_abs_sum": 0.0, + "verification_max_abs": 0.0, + } + + with torch.no_grad(): + actual_delta = torch.cat( + [ + target.data.view(-1).index_select(0, locations).float() - before + for target, locations, before, _ in verification + ] + ) + expected_delta = torch.cat([expected for _, _, _, expected in verification]) + difference = (actual_delta - expected_delta).abs() + exact_mismatches = actual_delta.ne(expected_delta) + mismatches = ~torch.isclose( + actual_delta, expected_delta, rtol=1e-6, atol=1e-8 + ) + stats = torch.stack( + ( + difference.sum(), + difference.max(), + exact_mismatches.sum().float(), + mismatches.sum().float(), + ) + ).cpu() + return { + "ok": True, + "verification_candidates": candidates, + "verification_samples": actual_delta.numel(), + "verification_exact_mismatches": int(stats[2]), + "verification_mismatches": int(stats[3]), + "verification_abs_sum": float(stats[0]), + "verification_max_abs": float(stats[1]), + } diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py new file mode 100644 index 00000000000..5d3050ca8e4 --- /dev/null +++ b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py @@ -0,0 +1,445 @@ +# 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. + +"""Remote sparse-refit receiver lifecycle for vLLM generation workers.""" + +import asyncio +import os +import tempfile +import threading +import time +from concurrent.futures import Future, ThreadPoolExecutor +from typing import Any, Literal, cast + +import uvicorn +from fastapi import FastAPI, Request +from fastapi.responses import JSONResponse + +from nemo_rl.distributed.virtual_cluster import ( + DEFAULT_GENERATION_PORT_RANGE_HIGH, + DEFAULT_GENERATION_PORT_RANGE_LOW, + _get_free_port_local, + _get_node_ip_local, +) +from nemo_rl.utils.weight_transfer_remote_sparse import ( + G_VLLM_REFIT_API_KEY_HEADER, + G_VLLM_REFIT_FLUSH_PATH, + G_VLLM_REFIT_S3_MANIFEST_PATH, + decode_sparse_payload, + download_s3_refit_payload, + merge_vllm_refit_receiver_timing, + refit_env_int, + vllm_refit_api_key, +) +from nemo_rl.utils.weight_transfer_zmq import ( + G_VLLM_REFIT_CHECKSUM_HEADER, + G_VLLM_REFIT_PAYLOAD_HEADER, + G_VLLM_REFIT_PRODUCER_HEADER, + G_VLLM_REFIT_TRANSFER_HEADER, + G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, + ZmqSparseRefitServer, +) + + +class VllmSparseRefitReceiver: + """Own the optional transport server, apply queue, and relay resources.""" + + def __init__(self, worker: Any) -> None: + self._worker = worker + self._refit_apply_queue_condition = threading.Condition() + self._refit_apply_executor = ThreadPoolExecutor( + max_workers=1, + thread_name_prefix="nrl-vllm-sparse-refit", + ) + self._refit_apply_futures: list[Future[dict[str, Any]]] = [] + self._refit_apply_pending_payloads: list[bytes] = [] + self._refit_seen_payloads: dict[tuple[str, int, int], str] = {} + self._refit_workers_share_node = False + self._refit_apply_queue_depth = refit_env_int( + "NRL_REFIT_APPLY_QUEUE_DEPTH", default=2 + ) + self._refit_apply_batch_size = refit_env_int( + "NRL_REFIT_APPLY_BATCH_SIZE", default=8 + ) + self._refit_batch_staging_dir = ( + os.getenv("NRL_REFIT_BATCH_STAGING_DIR") or "/dev/shm" + ) + self._refit_http_server: tuple[Any, threading.Thread, str] | None = None + self._zmq_refit_server: tuple[ZmqSparseRefitServer, str] | None = None + self._refit_async_loop: asyncio.AbstractEventLoop | None = None + + @property + def cfg(self) -> Any: + return self._worker.cfg + + @property + def llm(self) -> Any: + return self._worker.llm + + def set_worker_hostnames(self, hostnames: list[str]) -> None: + self._refit_workers_share_node = len(set(hostnames)) == 1 + + def start_sync_server(self) -> None: + llm = self.llm + if llm is None: + raise RuntimeError("vLLM is not initialized on this worker.") + self.set_worker_hostnames(llm.collective_rpc("report_node_hostname", args=())) + self._setup_vllm_refit_server() + + def shutdown(self) -> None: + self.stop_zmq_sparse_refit_relay() + if self._refit_http_server is not None: + self._refit_http_server[0].should_exit = True + + self._flush_queued_sparse_payloads() + self._refit_apply_executor.shutdown(wait=True) + + if self._refit_http_server is not None: + self._refit_http_server[1].join(timeout=5.0) + self._refit_http_server = None + + def _enqueue_sparse_payload_apply( + self, + payload: bytes, + payload_key: tuple[str, int, int], + checksum: str, + ) -> dict[str, Any]: + completed: list[Future[dict[str, Any]]] = [] + submitted = None + with self._refit_apply_queue_condition: + seen_checksum = self._refit_seen_payloads.get(payload_key) + if seen_checksum is not None: + if seen_checksum != checksum: + raise ValueError( + "A sparse refit payload ID was reused with different data." + ) + return {"ok": True, "payloads": 0, "duplicate": True} + while ( + len(self._refit_apply_futures) >= self._refit_apply_queue_depth + and not self._refit_apply_futures[0].done() + ): + self._refit_apply_queue_condition.wait() + while self._refit_apply_futures and self._refit_apply_futures[0].done(): + completed.append(self._refit_apply_futures.pop(0)) + response = self._collect_refit_apply_results(completed) + self._refit_seen_payloads[payload_key] = checksum + self._refit_apply_pending_payloads.append(payload) + if len(self._refit_apply_pending_payloads) == self._refit_apply_batch_size: + submitted = self._submit_pending_sparse_payloads() + if submitted is not None: + submitted.add_done_callback(self._notify_refit_apply_waiters) + return response + + def _submit_pending_sparse_payloads(self) -> Future[dict[str, Any]]: + payloads = tuple(self._refit_apply_pending_payloads) + self._refit_apply_pending_payloads.clear() + future = self._refit_apply_executor.submit( + self.update_weights_from_serialized_sparse_payloads, + payloads, + ) + self._refit_apply_futures.append(future) + return future + + def _notify_refit_apply_waiters(self, _future: Future[dict[str, Any]]) -> None: + with self._refit_apply_queue_condition: + self._refit_apply_queue_condition.notify_all() + + def _collect_refit_apply_results( + self, + futures: list[Future[dict[str, Any]]], + ) -> dict[str, Any]: + results = [future.result() for future in futures] + timing: dict[str, float] = {} + merge_vllm_refit_receiver_timing(timing, results, maximum=False) + return { + "ok": True, + "payloads": sum(int(result.get("payloads", 0)) for result in results), + **timing, + } + + @staticmethod + def _refit_collective_response(worker_results: Any) -> dict[str, Any]: + results = cast(list[dict[str, Any]], worker_results) + response = { + "ok": True, + **merge_vllm_refit_receiver_timing({}, results, maximum=True), + } + if any("verification_candidates" in result for result in results): + response["verification_candidates"] = max( + (int(result["verification_candidates"]) for result in results), + default=0, + ) + for key in ( + "verification_samples", + "verification_exact_mismatches", + "verification_mismatches", + "verification_abs_sum", + ): + response[key] = sum(result[key] for result in results) + response["verification_max_abs"] = max( + (float(result["verification_max_abs"]) for result in results), + default=0.0, + ) + return response + + def _refit_collective_rpc( + self, + method: str, + args: tuple[Any, ...], + ) -> Any: + llm = self.llm + if llm is None: + raise RuntimeError("vLLM is not initialized on this worker.") + if not self.cfg["vllm_cfg"]["async_engine"]: + return llm.collective_rpc(method, args=args) + if self._refit_async_loop is None: + raise RuntimeError("The async vLLM refit server loop is not initialized.") + return asyncio.run_coroutine_threadsafe( + llm.collective_rpc(method, args=args), + self._refit_async_loop, + ).result() + + def update_weights_from_serialized_sparse_payloads( + self, + serialized_payloads: tuple[bytes, ...], + ) -> dict[str, Any]: + """Apply a FIFO batch of sparse deltas through one collective RPC.""" + if self.llm is None: + raise RuntimeError("vLLM is not initialized on this worker.") + if not self._refit_workers_share_node: + results = [ + self._refit_collective_response( + self._refit_collective_rpc( + "update_weights_from_serialized_sparse_payload", + (payload,), + ) + ) + for payload in serialized_payloads + ] + timing: dict[str, float] = {} + merge_vllm_refit_receiver_timing(timing, results, maximum=False) + return {"ok": True, "payloads": len(serialized_payloads), **timing} + + with tempfile.TemporaryDirectory( + prefix="nemo_rl_refit_", dir=self._refit_batch_staging_dir + ) as staging_dir: + paths = [] + for index, payload in enumerate(serialized_payloads): + path = os.path.join(staging_dir, str(index)) + with open(path, "wb") as staged: + staged.write(payload) + paths.append(path) + try: + response = self._refit_collective_response( + self._refit_collective_rpc( + "update_weights_from_sparse_payload_files", + tuple(paths), + ) + ) + except Exception: + # Drain peers before TemporaryDirectory removes shared batch files. + self._refit_collective_rpc("synchronize_device", ()) + raise + response["payloads"] = len(serialized_payloads) + return response + + def _flush_queued_sparse_payloads(self) -> dict[str, Any]: + started = time.perf_counter() + submitted = None + with self._refit_apply_queue_condition: + if self._refit_apply_pending_payloads: + submitted = self._submit_pending_sparse_payloads() + futures = list(self._refit_apply_futures) + self._refit_apply_futures.clear() + self._refit_apply_queue_condition.notify_all() + payload_count = len(self._refit_seen_payloads) + batch_count = ( + payload_count + self._refit_apply_batch_size - 1 + ) // self._refit_apply_batch_size + if submitted is not None: + submitted.add_done_callback(self._notify_refit_apply_waiters) + response = self._collect_refit_apply_results(futures) + if futures: + assert self.llm is not None + response.update( + self._refit_collective_response( + self._refit_collective_rpc("finish_sparse_delta_refit", ()) + ) + ) + with self._refit_apply_queue_condition: + self._refit_seen_payloads.clear() + response.update( + payloads=payload_count, + batches=batch_count, + seconds=time.perf_counter() - started, + ) + if futures: + print( + "REFIT_RECEIVER_TIMING " + f"payloads={payload_count} batches={batch_count} " + f"total_s={response['seconds']:.3f} " + f"payload_total_s={response.get('receiver_total_s', 0.0):.3f} " + f"delta_verify_candidates=" + f"{response.get('verification_candidates', 0)} " + f"delta_verify_samples={response.get('verification_samples', 0)} " + f"delta_verify_exact_mismatches=" + f"{response.get('verification_exact_mismatches', 0)} " + f"delta_verify_mismatches=" + f"{response.get('verification_mismatches', 0)} " + f"delta_verify_max_abs=" + f"{response.get('verification_max_abs', 0.0):.8g}", + flush=True, + ) + return response + + async def _apply_s3_manifest_payload( + self, + manifest: dict[str, Any], + ) -> dict[str, Any]: + started = time.perf_counter() + body = await asyncio.to_thread(download_s3_refit_payload, manifest) + download_s = time.perf_counter() - started + key = str(manifest["key"]) + checksum = str(manifest["checksum"]) + result = await asyncio.to_thread( + self._enqueue_sparse_payload_apply, + body, + (key, -1, -1), + checksum, + ) + result["receiver_s3_download_s"] = download_s + return result + + async def _apply_zmq_payload(self, raw_request: Any) -> dict[str, Any]: + headers = raw_request.headers + transfer_id = headers.get(G_VLLM_REFIT_TRANSFER_HEADER, "") + producer_id = int(headers.get(G_VLLM_REFIT_PRODUCER_HEADER, "-1")) + payload_id = int(headers.get(G_VLLM_REFIT_PAYLOAD_HEADER, "-1")) + checksum = headers.get(G_VLLM_REFIT_CHECKSUM_HEADER, "") + if not transfer_id or producer_id < 0 or payload_id < 0 or not checksum: + raise ValueError("Missing or invalid ZeroMQ sparse refit payload headers.") + compressed = await raw_request.body() + started = time.perf_counter() + payload = await asyncio.to_thread( + decode_sparse_payload, + compressed, + checksum, + ) + decode_s = time.perf_counter() - started + result = await asyncio.to_thread( + self._enqueue_sparse_payload_apply, + payload, + (transfer_id, producer_id, payload_id), + checksum, + ) + result["receiver_zmq_decode_s"] = decode_s + return result + + def setup_api_server(self, app: Any) -> None: + token = vllm_refit_api_key( + self.cfg["vllm_cfg"].get("http_refit_api_key_env_var") + ) + + async def respond( + raw_request: Request, + action: Literal["s3", "flush", "zmq"], + ) -> JSONResponse: + if self.cfg["vllm_cfg"]["async_engine"]: + self._refit_async_loop = asyncio.get_running_loop() + if ( + token is not None + and raw_request.headers.get(G_VLLM_REFIT_API_KEY_HEADER) != token + ): + return JSONResponse( + content={"ok": False, "error": "unauthorized"}, status_code=403 + ) + try: + if action == "s3": + result = await self._apply_s3_manifest_payload( + await raw_request.json() + ) + elif action == "zmq": + result = await self._apply_zmq_payload(raw_request) + else: + result = await asyncio.to_thread(self._flush_queued_sparse_payloads) + except Exception as exc: + result = {"ok": False, "error": str(exc)} + return JSONResponse( + content=result, + status_code=200 if result.get("ok") is True else 500, + ) + + @app.post(G_VLLM_REFIT_S3_MANIFEST_PATH) + async def apply_s3_manifest_refit(raw_request: Request) -> JSONResponse: + return await respond(raw_request, "s3") + + @app.post(G_VLLM_REFIT_FLUSH_PATH) + async def flush_sparse_delta_refit(raw_request: Request) -> JSONResponse: + return await respond(raw_request, "flush") + + @app.post(G_VLLM_REFIT_ZMQ_PAYLOAD_PATH) + async def apply_zmq_sparse_refit(raw_request: Request) -> JSONResponse: + return await respond(raw_request, "zmq") + + def report_refit_server_base_url(self) -> str | None: + return self._refit_http_server[2] if self._refit_http_server else None + + def start_zmq_sparse_refit_relay(self, refit_urls: list[str]) -> str: + if self._zmq_refit_server is not None: + return self._zmq_refit_server[1] + port = self.cfg["vllm_cfg"].get( + "zmq_refit_server_port" + ) or _get_free_port_local( + self.cfg.get("port_range_low", DEFAULT_GENERATION_PORT_RANGE_LOW), + self.cfg.get("port_range_high", DEFAULT_GENERATION_PORT_RANGE_HIGH), + ) + server = ZmqSparseRefitServer( + refit_urls, + bind_address=f"tcp://0.0.0.0:{port}", + api_key_env_var=self.cfg["vllm_cfg"].get("http_refit_api_key_env_var"), + timeout_s=float(os.getenv("NRL_REFIT_ZMQ_TIMEOUT_S") or 600.0), + ) + server.start() + address = f"tcp://{_get_node_ip_local()}:{port}" + self._zmq_refit_server = (server, address) + print(f"Starting vLLM ZeroMQ refit relay on {address}", flush=True) + return address + + def stop_zmq_sparse_refit_relay(self) -> None: + if self._zmq_refit_server is not None: + self._zmq_refit_server[0].close() + self._zmq_refit_server = None + + def _setup_vllm_refit_server(self) -> None: + app = FastAPI() + self.setup_api_server(app) + port = self.cfg["vllm_cfg"].get( + "http_refit_server_port" + ) or _get_free_port_local( + self.cfg.get("port_range_low", DEFAULT_GENERATION_PORT_RANGE_LOW), + self.cfg.get("port_range_high", DEFAULT_GENERATION_PORT_RANGE_HIGH), + ) + server = uvicorn.Server( + uvicorn.Config( + app, + host="0.0.0.0", + port=port, + timeout_keep_alive=120, + ) + ) + thread = threading.Thread(target=server.run, daemon=True) + thread.start() + base_url = f"http://{_get_node_ip_local()}:{port}" + self._refit_http_server = (server, thread, base_url) + print(f"Starting vLLM refit server on {base_url}", flush=True) diff --git a/nemo_rl/models/generation/vllm/vllm_worker.py b/nemo_rl/models/generation/vllm/vllm_worker.py index 7559c5c3e5a..316ec5c4c64 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker.py +++ b/nemo_rl/models/generation/vllm/vllm_worker.py @@ -12,18 +12,12 @@ # See the License for the specific language governing permissions and # limitations under the License. -import asyncio import copy import gc import logging import os import sys -import tempfile -import threading -import time -import traceback -from concurrent.futures import Future, ThreadPoolExecutor -from typing import Any, Literal, Optional, cast +from typing import Any, Optional, cast import ray import torch @@ -31,12 +25,8 @@ from nemo_rl.distributed.batched_data_dict import BatchedDataDict from nemo_rl.distributed.virtual_cluster import ( - DEFAULT_GENERATION_PORT_RANGE_HIGH, - DEFAULT_GENERATION_PORT_RANGE_LOW, DEFAULT_VLLM_PORT_RANGE_LOW, DEFAULT_VLLM_PORTS_PER_ENGINE, - _get_free_port_local, - _get_node_ip_local, ) from nemo_rl.distributed.worker_group_utils import get_nsight_config_if_pattern_matches from nemo_rl.models.generation.interfaces import ( @@ -57,24 +47,6 @@ from nemo_rl.models.policy.utils import is_vllm_v1_engine_enabled from nemo_rl.utils.nsys import wrap_with_nvtx_name from nemo_rl.utils.nvml import log_gpu_memory_diagnostics -from nemo_rl.utils.weight_transfer_remote_sparse import ( - G_VLLM_REFIT_API_KEY_HEADER, - G_VLLM_REFIT_FLUSH_PATH, - G_VLLM_REFIT_S3_MANIFEST_PATH, - decode_sparse_payload, - download_s3_refit_payload, - merge_vllm_refit_receiver_timing, - refit_env_int, - vllm_refit_api_key, -) -from nemo_rl.utils.weight_transfer_zmq import ( - G_VLLM_REFIT_CHECKSUM_HEADER, - G_VLLM_REFIT_PAYLOAD_HEADER, - G_VLLM_REFIT_PRODUCER_HEADER, - G_VLLM_REFIT_TRANSFER_HEADER, - G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, - ZmqSparseRefitServer, -) logger = logging.getLogger(__name__) @@ -272,28 +244,17 @@ def __init__( if bundle_indices is not None and len(bundle_indices) == 1: bind_to_gpu_numa(int(ray.get_gpu_ids()[0])) - self._refit_apply_queue_condition = threading.Condition() - self._refit_apply_executor = ThreadPoolExecutor(max_workers=1) - self._refit_apply_futures: list[Future[dict[str, Any]]] = [] - self._refit_apply_pending_payloads: list[bytes] = [] - self._refit_seen_payloads: dict[tuple[str, int, int], str] = {} - self._refit_workers_share_node = False - self._refit_apply_queue_depth = refit_env_int( - "NRL_REFIT_APPLY_QUEUE_DEPTH", default=2 - ) - self._refit_apply_batch_size = refit_env_int( - "NRL_REFIT_APPLY_BATCH_SIZE", default=8 - ) - self._refit_batch_staging_dir = ( - os.getenv("NRL_REFIT_BATCH_STAGING_DIR") or "/dev/shm" - ) - self._refit_http_server: tuple[Any, threading.Thread, str] | None = None - self._zmq_refit_server: tuple[ZmqSparseRefitServer, str] | None = None - self._refit_async_loop: asyncio.AbstractEventLoop | None = None - self._init_config( config, bundle_indices, fraction_of_gpus, seed, extra_env_vars ) + self._sparse_refit_receiver: Any = None + if self.is_model_owner and self.cfg.get("refit_transport") is not None: + # Avoid receiver dependencies and threads for existing refit transports. + from nemo_rl.models.generation.vllm.vllm_sparse_refit import ( + VllmSparseRefitReceiver, + ) + + self._sparse_refit_receiver = VllmSparseRefitReceiver(self) if not self.is_model_owner: return @@ -697,310 +658,24 @@ def _get_raw_spec_counters(self) -> dict[str, float | list[float]]: metrics[metric.name] = metric.value return metrics - def _enqueue_sparse_payload_apply( - self, - payload: bytes, - payload_key: tuple[str, int, int], - checksum: str, - ) -> dict[str, Any]: - completed: list[Future[dict[str, Any]]] = [] - submitted = None - with self._refit_apply_queue_condition: - seen_checksum = self._refit_seen_payloads.get(payload_key) - if seen_checksum is not None: - if seen_checksum != checksum: - raise ValueError( - "A sparse refit payload ID was reused with different data." - ) - return {"ok": True, "payloads": 0, "duplicate": True} - while ( - len(self._refit_apply_futures) >= self._refit_apply_queue_depth - and not self._refit_apply_futures[0].done() - ): - self._refit_apply_queue_condition.wait() - while self._refit_apply_futures and self._refit_apply_futures[0].done(): - completed.append(self._refit_apply_futures.pop(0)) - response = self._collect_refit_apply_results(completed) - self._refit_seen_payloads[payload_key] = checksum - self._refit_apply_pending_payloads.append(payload) - if len(self._refit_apply_pending_payloads) == self._refit_apply_batch_size: - submitted = self._submit_pending_sparse_payloads() - if submitted is not None: - submitted.add_done_callback(self._notify_refit_apply_waiters) - return response - - def _submit_pending_sparse_payloads(self) -> Future[dict[str, Any]]: - payloads = tuple(self._refit_apply_pending_payloads) - self._refit_apply_pending_payloads.clear() - future = self._refit_apply_executor.submit( - self.update_weights_from_serialized_sparse_payloads, - payloads, - ) - self._refit_apply_futures.append(future) - return future - - def _notify_refit_apply_waiters(self, _future: Future[dict[str, Any]]) -> None: - with self._refit_apply_queue_condition: - self._refit_apply_queue_condition.notify_all() - - def _collect_refit_apply_results( - self, - futures: list[Future[dict[str, Any]]], - ) -> dict[str, Any]: - results = [future.result() for future in futures] - timing: dict[str, float] = {} - merge_vllm_refit_receiver_timing(timing, results, maximum=False) - return { - "ok": True, - "payloads": sum(int(result.get("payloads", 0)) for result in results), - **timing, - } - - @staticmethod - def _refit_collective_response(worker_results: Any) -> dict[str, Any]: - results = cast(list[dict[str, Any]], worker_results) - response = { - "ok": True, - **merge_vllm_refit_receiver_timing({}, results, maximum=True), - } - if any("verification_candidates" in result for result in results): - response["verification_candidates"] = max( - (int(result["verification_candidates"]) for result in results), - default=0, - ) - for key in ( - "verification_samples", - "verification_exact_mismatches", - "verification_mismatches", - "verification_abs_sum", - ): - response[key] = sum(result[key] for result in results) - response["verification_max_abs"] = max( - (float(result["verification_max_abs"]) for result in results), - default=0.0, - ) - return response - - def _refit_collective_rpc( - self, - method: str, - args: tuple[Any, ...], - ) -> Any: - return self.llm.collective_rpc(method, args=args) - - def update_weights_from_serialized_sparse_payloads( - self, - serialized_payloads: tuple[bytes, ...], - ) -> dict[str, Any]: - """Apply a FIFO batch of sparse deltas through one collective RPC.""" - if self.llm is None: - raise RuntimeError("vLLM is not initialized on this worker.") - if not self._refit_workers_share_node: - results = [ - self._refit_collective_response( - self._refit_collective_rpc( - "update_weights_from_serialized_sparse_payload", - (payload,), - ) - ) - for payload in serialized_payloads - ] - timing: dict[str, float] = {} - merge_vllm_refit_receiver_timing(timing, results, maximum=False) - return {"ok": True, "payloads": len(serialized_payloads), **timing} - - with tempfile.TemporaryDirectory( - prefix="nemo_rl_refit_", dir=self._refit_batch_staging_dir - ) as staging_dir: - paths = [] - for index, payload in enumerate(serialized_payloads): - path = os.path.join(staging_dir, str(index)) - with open(path, "wb") as staged: - staged.write(payload) - paths.append(path) - try: - response = self._refit_collective_response( - self._refit_collective_rpc( - "update_weights_from_sparse_payload_files", - tuple(paths), - ) - ) - except Exception: - # Drain peers before TemporaryDirectory removes shared batch files. - self._refit_collective_rpc("synchronize_device", ()) - raise - response["payloads"] = len(serialized_payloads) - return response - - def _flush_queued_sparse_payloads(self) -> dict[str, Any]: - started = time.perf_counter() - submitted = None - with self._refit_apply_queue_condition: - if self._refit_apply_pending_payloads: - submitted = self._submit_pending_sparse_payloads() - futures = list(self._refit_apply_futures) - self._refit_apply_futures.clear() - self._refit_apply_queue_condition.notify_all() - payload_count = len(self._refit_seen_payloads) - batch_count = ( - payload_count + self._refit_apply_batch_size - 1 - ) // self._refit_apply_batch_size - if submitted is not None: - submitted.add_done_callback(self._notify_refit_apply_waiters) - response = self._collect_refit_apply_results(futures) - if futures: - assert self.llm is not None - response.update( - self._refit_collective_response( - self._refit_collective_rpc("finish_sparse_delta_refit", ()) - ) - ) - with self._refit_apply_queue_condition: - self._refit_seen_payloads.clear() - response.update( - payloads=payload_count, - batches=batch_count, - seconds=time.perf_counter() - started, - ) - if futures: - print( - "REFIT_RECEIVER_TIMING " - f"payloads={payload_count} batches={batch_count} " - f"total_s={response['seconds']:.3f} " - f"payload_total_s={response.get('receiver_total_s', 0.0):.3f} " - f"delta_verify_candidates=" - f"{response.get('verification_candidates', 0)} " - f"delta_verify_samples={response.get('verification_samples', 0)} " - f"delta_verify_exact_mismatches=" - f"{response.get('verification_exact_mismatches', 0)} " - f"delta_verify_mismatches=" - f"{response.get('verification_mismatches', 0)} " - f"delta_verify_max_abs=" - f"{response.get('verification_max_abs', 0.0):.8g}", - flush=True, - ) - return response - - async def _apply_s3_manifest_payload( - self, - manifest: dict[str, Any], - ) -> dict[str, Any]: - started = time.perf_counter() - body = await asyncio.to_thread(download_s3_refit_payload, manifest) - download_s = time.perf_counter() - started - key = str(manifest["key"]) - checksum = str(manifest["checksum"]) - result = await asyncio.to_thread( - self._enqueue_sparse_payload_apply, - body, - (key, -1, -1), - checksum, - ) - result["receiver_s3_download_s"] = download_s - return result - - async def _apply_zmq_payload(self, raw_request: Any) -> dict[str, Any]: - headers = raw_request.headers - transfer_id = headers.get(G_VLLM_REFIT_TRANSFER_HEADER, "") - producer_id = int(headers.get(G_VLLM_REFIT_PRODUCER_HEADER, "-1")) - payload_id = int(headers.get(G_VLLM_REFIT_PAYLOAD_HEADER, "-1")) - checksum = headers.get(G_VLLM_REFIT_CHECKSUM_HEADER, "") - if not transfer_id or producer_id < 0 or payload_id < 0 or not checksum: - raise ValueError("Missing or invalid ZeroMQ sparse refit payload headers.") - compressed = await raw_request.body() - started = time.perf_counter() - payload = await asyncio.to_thread( - decode_sparse_payload, - compressed, - checksum, - ) - decode_s = time.perf_counter() - started - result = await asyncio.to_thread( - self._enqueue_sparse_payload_apply, - payload, - (transfer_id, producer_id, payload_id), - checksum, - ) - result["receiver_zmq_decode_s"] = decode_s - return result - - def _setup_vllm_refit_api_server(self, app: Any) -> None: - from fastapi import Request - from fastapi.responses import JSONResponse - - token = vllm_refit_api_key( - self.cfg["vllm_cfg"].get("http_refit_api_key_env_var") - ) - - async def respond( - raw_request: Request, - action: Literal["s3", "flush", "zmq"], - ) -> JSONResponse: - if self.cfg["vllm_cfg"]["async_engine"]: - self._refit_async_loop = asyncio.get_running_loop() - if ( - token is not None - and raw_request.headers.get(G_VLLM_REFIT_API_KEY_HEADER) != token - ): - return JSONResponse( - content={"ok": False, "error": "unauthorized"}, status_code=403 - ) - try: - if action == "s3": - result = await self._apply_s3_manifest_payload( - await raw_request.json() - ) - elif action == "zmq": - result = await self._apply_zmq_payload(raw_request) - else: - result = await asyncio.to_thread(self._flush_queued_sparse_payloads) - except Exception as exc: - result = {"ok": False, "error": str(exc)} - return JSONResponse( - content=result, - status_code=200 if result.get("ok") is True else 500, - ) - - @app.post(G_VLLM_REFIT_S3_MANIFEST_PATH) - async def apply_s3_manifest_refit(raw_request: Request) -> JSONResponse: - return await respond(raw_request, "s3") - - @app.post(G_VLLM_REFIT_FLUSH_PATH) - async def flush_sparse_delta_refit(raw_request: Request) -> JSONResponse: - return await respond(raw_request, "flush") - - @app.post(G_VLLM_REFIT_ZMQ_PAYLOAD_PATH) - async def apply_zmq_sparse_refit(raw_request: Request) -> JSONResponse: - return await respond(raw_request, "zmq") + def _require_sparse_refit_receiver(self) -> Any: + if self._sparse_refit_receiver is None: + raise RuntimeError("Remote sparse refit is not enabled for this worker.") + return self._sparse_refit_receiver def report_refit_server_base_url(self) -> str | None: - return self._refit_http_server[2] if self._refit_http_server else None + receiver = self._sparse_refit_receiver + return receiver.report_refit_server_base_url() if receiver is not None else None def start_zmq_sparse_refit_relay(self, refit_urls: list[str]) -> str: - if self._zmq_refit_server is not None: - return self._zmq_refit_server[1] - port = self.cfg["vllm_cfg"].get( - "zmq_refit_server_port" - ) or _get_free_port_local( - self.cfg.get("port_range_low", DEFAULT_GENERATION_PORT_RANGE_LOW), - self.cfg.get("port_range_high", DEFAULT_GENERATION_PORT_RANGE_HIGH), - ) - server = ZmqSparseRefitServer( - refit_urls, - bind_address=f"tcp://0.0.0.0:{port}", - api_key_env_var=self.cfg["vllm_cfg"].get("http_refit_api_key_env_var"), - timeout_s=float(os.getenv("NRL_REFIT_ZMQ_TIMEOUT_S") or 600.0), + return self._require_sparse_refit_receiver().start_zmq_sparse_refit_relay( + refit_urls ) - server.start() - address = f"tcp://{_get_node_ip_local()}:{port}" - self._zmq_refit_server = (server, address) - print(f"Starting vLLM ZeroMQ refit relay on {address}", flush=True) - return address def stop_zmq_sparse_refit_relay(self) -> None: - if self._zmq_refit_server is not None: - self._zmq_refit_server[0].close() - self._zmq_refit_server = None + receiver = self._sparse_refit_receiver + if receiver is not None: + receiver.stop_zmq_sparse_refit_relay() class VllmGenerationWorkerImpl(BaseVllmGenerationWorker): @@ -1017,37 +692,8 @@ def post_init(self): self.llm.collective_rpc( "load_mtp_weights_from_disk", args=(self.model_name,) ) - if self.cfg.get("refit_transport") is not None: - self._refit_workers_share_node = ( - len(set(self.llm.collective_rpc("report_node_hostname", args=()))) == 1 - ) - self._setup_vllm_refit_server() - - def _setup_vllm_refit_server(self) -> None: - import uvicorn - from fastapi import FastAPI - - app = FastAPI() - self._setup_vllm_refit_api_server(app) - port = self.cfg["vllm_cfg"].get( - "http_refit_server_port" - ) or _get_free_port_local( - self.cfg.get("port_range_low", DEFAULT_GENERATION_PORT_RANGE_LOW), - self.cfg.get("port_range_high", DEFAULT_GENERATION_PORT_RANGE_HIGH), - ) - server = uvicorn.Server( - uvicorn.Config( - app, - host="0.0.0.0", - port=port, - timeout_keep_alive=120, - ) - ) - thread = threading.Thread(target=server.run, daemon=True) - thread.start() - base_url = f"http://{_get_node_ip_local()}:{port}" - self._refit_http_server = (server, thread, base_url) - print(f"Starting vLLM refit server on {base_url}", flush=True) + if self._sparse_refit_receiver is not None: + self._sparse_refit_receiver.start_sync_server() def init_collective( self, @@ -1173,6 +819,8 @@ def generate( 1 ].logprob except Exception: + import traceback + traceback.print_exc() logprobs_list.append(full_logprobs) @@ -1370,6 +1018,8 @@ def update_weights_via_ipc_zmq(self) -> bool: return True except Exception as e: print(f"Exception during collective_rpc for weight update: {e}") + import traceback + traceback.print_exc() return False @@ -1399,6 +1049,8 @@ def update_weights_from_collective(self) -> bool: return True except Exception as e: print(f"Exception during collective_rpc for weight update: {e}") + import traceback + traceback.print_exc() return False @@ -1467,16 +1119,8 @@ def wake_up(self, **kwargs): def shutdown(self) -> bool: """Clean up vLLM resources.""" try: - self.stop_zmq_sparse_refit_relay() - if self._refit_http_server is not None: - self._refit_http_server[0].should_exit = True - - self._flush_queued_sparse_payloads() - self._refit_apply_executor.shutdown(wait=True) - - if self._refit_http_server is not None: - self._refit_http_server[1].join(timeout=5.0) - self._refit_http_server = None + if self._sparse_refit_receiver is not None: + self._sparse_refit_receiver.shutdown() if self.llm is not None: # Clean up extension resources (e.g., ZMQ sockets) diff --git a/nemo_rl/models/generation/vllm/vllm_worker_async.py b/nemo_rl/models/generation/vllm/vllm_worker_async.py index f9baf78fc4d..17fcd30e489 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker_async.py +++ b/nemo_rl/models/generation/vllm/vllm_worker_async.py @@ -218,17 +218,6 @@ def __init__( self.llm = None self.vllm_device_ids = None - def _refit_collective_rpc( - self, - method: str, - args: tuple[Any, ...], - ) -> Any: - if self._refit_async_loop is None: - raise RuntimeError("The async vLLM refit server loop is not initialized.") - return asyncio.run_coroutine_threadsafe( - self.llm.collective_rpc(method, args=args), self._refit_async_loop - ).result() - def _return_routed_experts_enabled(self) -> bool: engine_args = getattr(self, "llm_async_engine_args", None) if bool(getattr(engine_args, "enable_return_routed_experts", False)): @@ -440,9 +429,9 @@ async def post_init_async(self): await self.llm.collective_rpc( "load_mtp_weights_from_disk", args=(self.model_name,) ) - if self.cfg.get("refit_transport") is not None: + if self._sparse_refit_receiver is not None: hostnames = await self.llm.collective_rpc("report_node_hostname", args=()) - self._refit_workers_share_node = len(set(hostnames)) == 1 + self._sparse_refit_receiver.set_worker_hostnames(hostnames) async def get_reserved_url(self) -> Optional[str]: """Return the URL from the reserved socket, available before model loading.""" @@ -930,8 +919,8 @@ def _setup_vllm_server(self) -> "tuple[threading.Thread, str, uvicorn.Server]": app = FastAPI() app = self._setup_vllm_openai_api_server(app) - if self.cfg.get("refit_transport") is not None: - self._setup_vllm_refit_api_server(app) + if self._sparse_refit_receiver is not None: + self._sparse_refit_receiver.setup_api_server(app) ######################################## # Server spinup @@ -1556,9 +1545,8 @@ async def wake_up_async(self, **kwargs): async def shutdown(self) -> bool: """Clean up vLLM resources.""" try: - self.stop_zmq_sparse_refit_relay() - await asyncio.to_thread(self._flush_queued_sparse_payloads) - self._refit_apply_executor.shutdown(wait=True) + if self._sparse_refit_receiver is not None: + await asyncio.to_thread(self._sparse_refit_receiver.shutdown) if self.llm is not None: # Clean up extension resources (e.g., ZMQ sockets) diff --git a/nemo_rl/models/policy/lm_policy.py b/nemo_rl/models/policy/lm_policy.py index 226641b2039..397b4e086b5 100644 --- a/nemo_rl/models/policy/lm_policy.py +++ b/nemo_rl/models/policy/lm_policy.py @@ -12,7 +12,6 @@ # See the License for the specific language governing permissions and # limitations under the License. import os -import sys import warnings from collections import defaultdict from contextlib import nullcontext @@ -941,13 +940,6 @@ 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 init_remote_sparse_delta_baseline(self, transport: str) -> list[ray.ObjectRef]: - """Initialize source-side sparse-delta baselines for remote refit.""" - return self._run_remote_sparse_refit_workers( - "init_remote_sparse_delta_baseline", - transport=transport, - ) - def finish_inference(self) -> None: """Offload policy model to CPU after inference.""" futures = self.worker_group.run_all_workers_single_data("finish_inference") @@ -1071,45 +1063,6 @@ def set_rollout_num_gpus_per_engine(self, num_gpus_per_engine: int) -> None: ) ) - def stream_remote_sparse_weights( - self, - transport: str, - targets: list[str], - *, - transfer_id: str, - api_key_env_var: Optional[str], - timeout_s: float, - ) -> list[ray.ObjectRef]: - """Stream sparse deltas through the selected remote value plane.""" - return self._run_remote_sparse_refit_workers( - "stream_remote_sparse_weights", - transport=transport, - targets=targets, - transfer_id=transfer_id, - api_key_env_var=api_key_env_var, - timeout_s=timeout_s, - ) - - def finish_remote_sparse_delta_sync(self, succeeded: bool) -> list[ray.ObjectRef]: - return self.worker_group.run_all_workers_single_data( - "finish_remote_sparse_delta_sync", succeeded=succeeded - ) - - def _run_remote_sparse_refit_workers( - self, - method_name: str, - **common_kwargs: Any, - ) -> list[ray.ObjectRef]: - worker_count = len(self.worker_group.workers) - return self.worker_group.run_all_workers_multiple_data( - method_name, - common_kwargs={ - **common_kwargs, - "shard_count": worker_count, - }, - shard_rank=list(range(worker_count)), - ) - def broadcast_weights_for_collective( self, kv_scales: Optional[dict[str, float]] = None ) -> list[ray.ObjectRef]: @@ -1199,7 +1152,7 @@ def __del__(self) -> None: the object is lost due to leaving a function scope. It's always recommended that the user calls worker_group.shutdown(). """ - if not sys.is_finalizing() and hasattr(self, "worker_group"): + if hasattr(self, "worker_group"): self.worker_group.shutdown(cleanup_method="shutdown") def start_gpu_profiling(self) -> None: diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 7cef71bf17d..32f284fe740 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -64,6 +64,7 @@ ) from nemo_rl.models.megatron.pipeline_parallel import ( broadcast_loss_metrics_from_last_stage, + broadcast_obj_from_pp_rank, broadcast_tensors_from_last_stage, ) from nemo_rl.models.megatron.router_replay import router_replay_enabled @@ -98,13 +99,6 @@ from nemo_rl.utils.packed_tensor import packed_broadcast_producer from nemo_rl.utils.r3_trace import maybe_r3_trace_stage from nemo_rl.utils.timer import Timer -from nemo_rl.utils.weight_transfer_remote_sparse import ( - SparseDeltaStreamResult, - init_sparse_delta_baseline_from_iterator, - stream_sparse_delta_payloads_via_s3_manifest, -) -from nemo_rl.utils.weight_transfer_sparse_codec import DeltaCompressionTracker -from nemo_rl.utils.weight_transfer_zmq import stream_sparse_delta_payloads_via_zmq TokenizerType = TypeVar("TokenizerType", bound=PreTrainedTokenizerBase) @@ -374,9 +368,14 @@ def __init__( delta_config = None if generation_config is not None and generation_config["backend"] == "vllm": delta_config = cast(VllmConfig, generation_config).get("delta_compression") - self.delta_weight_transfer_tracker = ( - DeltaCompressionTracker(delta_config) if delta_config else None - ) + self._remote_sparse_refit = None + if delta_config: + # Keep codec and remote transport state out of standard policy workers. + from nemo_rl.models.policy.workers.megatron_remote_sparse_refit import ( + MegatronRemoteSparseRefit, + ) + + self._remote_sparse_refit = MegatronRemoteSparseRefit(self, delta_config) self.defer_fp32_logits = self.cfg["megatron_cfg"].get( "defer_fp32_logits", None @@ -1817,12 +1816,7 @@ def init_remote_sparse_delta_baseline( shard_count: int, transport: str, ) -> None: - """Initialize the source-side baseline for remote sparse refit.""" - tracker = self.delta_weight_transfer_tracker - assert tracker is not None - init_sparse_delta_baseline_from_iterator( - self._iter_params_with_optional_kv_scales(), - delta_tracker=tracker, + self._require_remote_sparse_refit().initialize_baseline( shard_rank=shard_rank, shard_count=shard_count, transport=transport, @@ -1840,34 +1834,21 @@ def stream_remote_sparse_weights( timeout_s: float, shard_rank: int, shard_count: int, - ) -> SparseDeltaStreamResult: - """Stream compressed sparse deltas through the selected value plane.""" - tracker = self.delta_weight_transfer_tracker - assert tracker is not None - streamer = { - "s3": stream_sparse_delta_payloads_via_s3_manifest, - "zmq": stream_sparse_delta_payloads_via_zmq, - }.get(transport) - if streamer is None: - raise ValueError( - f"Unsupported remote sparse refit transport {transport!r}." - ) - result = streamer( - self._iter_params_with_optional_kv_scales(), - delta_tracker=tracker, - refit_targets=targets, + ) -> dict[str, int]: + return self._require_remote_sparse_refit().stream( + transport, + targets, transfer_id=transfer_id, api_key_env_var=api_key_env_var, timeout_s=timeout_s, shard_rank=shard_rank, shard_count=shard_count, ) - # HF export can enqueue work on auxiliary streams. Drain the device before - # this actor returns so every policy rank resumes training from the same - # completed CUDA boundary. - if torch.cuda.is_available(): - torch.cuda.synchronize() - return result + + def _require_remote_sparse_refit(self) -> Any: + if self._remote_sparse_refit is None: + raise RuntimeError("Remote sparse refit is not enabled for this worker.") + return self._remote_sparse_refit def _get_refit_conversion_tasks(self) -> list[Any]: if self.refit_conversion_tasks is None: @@ -1876,13 +1857,7 @@ def _get_refit_conversion_tasks(self) -> list[Any]: return self.refit_conversion_tasks def finish_remote_sparse_delta_sync(self, *, succeeded: bool) -> None: - tracker = self.delta_weight_transfer_tracker - if tracker is None: - raise RuntimeError("Sparse delta tracker is not initialized.") - if succeeded: - tracker.on_sync_succeeded() - else: - tracker.on_sync_failed() + self._require_remote_sparse_refit().finish(succeeded) def _calculate_refit_param_info(self) -> list[tuple[str, int]]: """Calculate parameter information for refit. @@ -1901,7 +1876,7 @@ def _calculate_refit_param_info(self) -> list[tuple[str, int]]: conversion_tasks = self._get_refit_conversion_tasks() param_info = [] - def calculate_size_in_bytes(param, mapping): + def calculate_size_in_bytes(param, tp_size, ep_size): if param is None: # need to broadcast for other pp ranks size_in_bytes = None @@ -1916,15 +1891,11 @@ def calculate_size_in_bytes(param, mapping): } scale = prec_to_bytes[self.dtype] / prec_to_bytes[param.dtype] size_in_bytes = ( - param.element_size() - * param.numel() - * mapping.tp_size - * (mapping.ep_size if mapping.is_expert else 1) - * scale + param.element_size() * param.numel() * tp_size * ep_size * scale ) - # Match Megatron Bridge export semantics for tied or replicated weights. - return mapping.broadcast_obj_from_pp_rank(size_in_bytes) + # Broadcast size_in_bytes across pipeline parallel ranks + return broadcast_obj_from_pp_rank(size_in_bytes) for task in conversion_tasks: param_info.append( @@ -1932,7 +1903,8 @@ def calculate_size_in_bytes(param, mapping): task.param_name, calculate_size_in_bytes( task.param_weight, - task.mapping, + task.mapping.tp_size, + task.mapping.ep_size if task.mapping.is_expert else 1, ), ) ) diff --git a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py new file mode 100644 index 00000000000..895f143999c --- /dev/null +++ b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py @@ -0,0 +1,87 @@ +# 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. + +"""Optional remote sparse-refit state owned by a Megatron policy worker.""" + +from typing import Any + +import torch + +from nemo_rl.utils.weight_transfer_remote_sparse import ( + SparseDeltaStreamResult, + init_sparse_delta_baseline_from_iterator, + stream_sparse_delta_payloads_via_s3_manifest, +) +from nemo_rl.utils.weight_transfer_sparse_codec import DeltaCompressionTracker +from nemo_rl.utils.weight_transfer_zmq import stream_sparse_delta_payloads_via_zmq + + +class MegatronRemoteSparseRefit: + def __init__(self, worker: Any, delta_config: dict[str, Any]) -> None: + self._worker = worker + self._tracker = DeltaCompressionTracker(delta_config) + + def initialize_baseline( + self, + *, + shard_rank: int, + shard_count: int, + transport: str, + ) -> None: + init_sparse_delta_baseline_from_iterator( + self._worker._iter_params_with_optional_kv_scales(), + delta_tracker=self._tracker, + shard_rank=shard_rank, + shard_count=shard_count, + transport=transport, + ) + + def stream( + self, + transport: str, + targets: list[str], + *, + transfer_id: str, + api_key_env_var: str | None, + timeout_s: float, + shard_rank: int, + shard_count: int, + ) -> SparseDeltaStreamResult: + streamer = { + "s3": stream_sparse_delta_payloads_via_s3_manifest, + "zmq": stream_sparse_delta_payloads_via_zmq, + }.get(transport) + if streamer is None: + raise ValueError( + f"Unsupported remote sparse refit transport {transport!r}." + ) + result = streamer( + self._worker._iter_params_with_optional_kv_scales(), + delta_tracker=self._tracker, + refit_targets=targets, + transfer_id=transfer_id, + api_key_env_var=api_key_env_var, + timeout_s=timeout_s, + shard_rank=shard_rank, + shard_count=shard_count, + ) + if torch.cuda.is_available(): + torch.cuda.synchronize() + return result + + def finish(self, succeeded: bool) -> None: + if succeeded: + self._tracker.on_sync_succeeded() + else: + self._tracker.on_sync_failed() diff --git a/nemo_rl/utils/weight_transfer_remote_sparse.py b/nemo_rl/utils/weight_transfer_remote_sparse.py index f839c1d8dcf..e2be67a8839 100644 --- a/nemo_rl/utils/weight_transfer_remote_sparse.py +++ b/nemo_rl/utils/weight_transfer_remote_sparse.py @@ -54,6 +54,7 @@ class SparseDeltaStreamResult(TypedDict): @cache def _s3_client(region: str) -> Any: + # Keep the AWS runtime unloaded for ZeroMQ-only jobs. from awscrt.auth import AwsCredentialsProvider from awscrt.io import ClientBootstrap, DefaultHostResolver, EventLoopGroup from awscrt.s3 import S3Client, create_default_s3_signing_config @@ -79,6 +80,7 @@ def _s3_client(region: str) -> Any: class _S3ObjectStore: def __init__(self, *, bucket: str, region: str) -> None: + # Keep the AWS runtime unloaded for ZeroMQ-only jobs. from awscrt.s3 import S3RequestType self.bucket = bucket @@ -93,6 +95,7 @@ def put(self, key: str, body: bytes) -> None: ).finished_future.result() def get(self, key: str) -> bytes: + # Keep the AWS runtime unloaded for ZeroMQ-only jobs. from awscrt.http import HttpHeaders body = bytearray() @@ -129,6 +132,7 @@ def delete(self, key: str) -> None: ).finished_future.result() def _request(self, method: str, key: str, body: bytes | None = None) -> Any: + # Keep the AWS runtime unloaded for ZeroMQ-only jobs. from awscrt.http import HttpHeaders, HttpRequest headers = HttpHeaders( diff --git a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py index f89adfa0267..eab7fab6ad9 100644 --- a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py +++ b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py @@ -26,6 +26,33 @@ from nemo_rl.weight_sync.interfaces import WeightSynchronizer +def validate_vllm_remote_sparse_refit( + config: Any, + *, + colocated: bool, + megatron_enabled: bool, +) -> str | None: + """Validate the optional transport without exposing its rules to GRPO.""" + transport = config.get("refit_transport") + if transport not in (None, "vllm_s3_sparse", "vllm_zmq_sparse"): + raise ValueError(f"Unsupported vLLM refit transport {transport!r}.") + vllm_cfg = config["vllm_cfg"] + if transport is not None and ( + colocated + or not megatron_enabled + or vllm_cfg["precision"] == "fp8" + or vllm_cfg["kv_cache_dtype"].startswith("fp8") + or not config.get("delta_compression") + or config.get("quant_cfg") + or config.get("real_quant") + ): + raise ValueError( + f"{transport} requires a non-colocated Megatron policy, BF16/FP16 " + "vLLM, delta compression, and an unquantized rollout." + ) + return transport + + class VllmRemoteSparseWeightSynchronizer(WeightSynchronizer): def __init__( self, @@ -74,9 +101,10 @@ def sync_weights( succeeded = False try: transfer_id = uuid.uuid4().hex - refs = self._policy.stream_remote_sparse_weights( - self._transport, - self._targets, + refs = self._run_policy_workers( + "stream_remote_sparse_weights", + transport=self._transport, + targets=self._targets, transfer_id=transfer_id, api_key_env_var=self._api_key_env_var, timeout_s=self._request_timeout_s, @@ -166,7 +194,10 @@ def sync_weights( timeout_s=min(self._request_timeout_s, 60.0), ) self._baseline_commit_refs = ( - self._policy.finish_remote_sparse_delta_sync(succeeded) + self._policy.worker_group.run_all_workers_single_data( + "finish_remote_sparse_delta_sync", + succeeded=succeeded, + ) ) self._stale = False return { @@ -191,16 +222,46 @@ def is_stale(self) -> bool: def mark_stale(self) -> None: self._stale = True + def _run_policy_workers(self, method_name: str, **kwargs: Any) -> list[Any]: + worker_group = self._policy.worker_group + worker_count = len(worker_group.workers) + return worker_group.run_all_workers_multiple_data( + method_name, + common_kwargs={**kwargs, "shard_count": worker_count}, + shard_rank=list(range(worker_count)), + ) + + def _run_generation_workers(self, method_name: str, **kwargs: Any) -> list[Any]: + worker_group = self._generation.worker_group + if not worker_group or not worker_group.workers: + raise RuntimeError("vLLM worker group is not initialized.") + return ray.get( + worker_group.run_all_workers_single_data( + method_name, + run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], + **kwargs, + ) + ) + def init_communicator(self) -> None: - self._baseline_init_refs = self._policy.init_remote_sparse_delta_baseline( - self._transport + self._baseline_init_refs = self._run_policy_workers( + "init_remote_sparse_delta_baseline", + transport=self._transport, ) - self._refit_urls = self._generation.report_refit_server_base_urls() + self._refit_urls = [ + url + for url in self._run_generation_workers("report_refit_server_base_url") + if url + ] self._targets = self._refit_urls if self._transport == "zmq": - self._targets = self._generation.start_zmq_sparse_refit_relays( - self._refit_urls - ) + self._targets = [ + address + for address in self._run_generation_workers( + "start_zmq_sparse_refit_relay", refit_urls=self._refit_urls + ) + if address + ] if not self._refit_urls or not self._targets: raise ValueError( f"vLLM {self._transport} sparse refit endpoints are missing." @@ -213,7 +274,7 @@ def shutdown(self) -> None: ): ray.cancel(ref, force=False) if self._transport == "zmq": - self._generation.stop_zmq_sparse_refit_relays() + self._run_generation_workers("stop_zmq_sparse_refit_relay") self._baseline_init_refs = None self._baseline_commit_refs = None self._refit_urls = [] diff --git a/pyrefly.toml b/pyrefly.toml index 3aece31741a..4ee30d2634b 100644 --- a/pyrefly.toml +++ b/pyrefly.toml @@ -201,6 +201,9 @@ project-includes = [ "nemo_rl/weight_sync/interfaces.py", "nemo_rl/weight_sync/ipc_weight_synchronizer.py", "nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py", + "nemo_rl/models/generation/vllm/vllm_sparse_delta.py", + "nemo_rl/models/generation/vllm/vllm_sparse_refit.py", + "nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py", "nemo_rl/utils/weight_transfer_remote_sparse.py", "nemo_rl/utils/weight_transfer_sparse_codec.py", "nemo_rl/utils/weight_transfer_zmq.py", diff --git a/tests/unit/algorithms/test_grpo.py b/tests/unit/algorithms/test_grpo.py index bc8833f1fd0..5e26f583218 100644 --- a/tests/unit/algorithms/test_grpo.py +++ b/tests/unit/algorithms/test_grpo.py @@ -1891,25 +1891,12 @@ def fake_batched_message_log_to_flat_message(*_args, **_kwargs): lambda *_args, **_kwargs: seq_logprob_error_result, ) - dynamic_sampling_calls = 0 - - def fake_dynamic_sampling(repeated_batch, *_args, **_kwargs): - nonlocal dynamic_sampling_calls - dynamic_sampling_calls += 1 - repeated_batch["filtered_reward"] = repeated_batch["total_reward"] - repeated_batch["baseline"] = torch.zeros(repeated_batch.size) - repeated_batch["std"] = torch.ones(repeated_batch.size) - complete = dynamic_sampling_calls == 2 - return repeated_batch, complete, None if complete else repeated_batch, {} - - monkeypatch.setattr(grpo_mod, "dynamic_sampling", fake_dynamic_sampling) - master_config = mock_grpo_components["master_config"] master_config.grpo["max_num_steps"] = 1 master_config.grpo["max_num_epochs"] = 1 master_config.grpo["val_period"] = 0 master_config.grpo["val_at_start"] = False - master_config.grpo["use_dynamic_sampling"] = True + master_config.grpo["use_dynamic_sampling"] = False grpo_mod.grpo_train( mock_grpo_components["policy"], @@ -1940,7 +1927,6 @@ def fake_dynamic_sampling(repeated_batch, *_args, **_kwargs): assert train_metrics["min_seq_mult_prob_error_after_mask"] == 1.0 assert train_metrics["num_masked_seqs_by_logprob_error"] == 2 assert train_metrics["masked_correct_pct"] == 0.5 - assert dynamic_sampling_calls == 2 assert any( call.args[0] == {"delta/changed_pct": 4.0} and call.kwargs.get("prefix") == "refit" diff --git a/tests/unit/models/generation/test_vllm_backend.py b/tests/unit/models/generation/test_vllm_backend.py index 0d29ed7d60e..ea667f6c86c 100644 --- a/tests/unit/models/generation/test_vllm_backend.py +++ b/tests/unit/models/generation/test_vllm_backend.py @@ -18,16 +18,13 @@ import contextlib import json -from types import MethodType, SimpleNamespace -from typing import Any +from types import SimpleNamespace from unittest.mock import MagicMock import pytest import torch from safetensors.torch import save_file -from nemo_rl.utils.weight_transfer_sparse_codec import encode_sparse_infos - def _make_collective_update_extension(backend): ext = backend.VllmInternalWorkerExtension.__new__( @@ -36,7 +33,7 @@ def _make_collective_update_extension(backend): state_info = object() ext.state_dict_info = {"model.weight": state_info} ext.model_update_group = object() - ext.model_runner = SimpleNamespace(model=object(), vllm_config=object()) + ext.model_runner = SimpleNamespace(model=object()) ext.model_config = object() ext.device = object() return ext, state_info @@ -90,297 +87,14 @@ def _patch_vllm_postload(monkeypatch): return process_weights -def _attach_tensor_attrs(tensor: torch.Tensor, **attrs: object) -> torch.Tensor: - for name, value in attrs.items(): - setattr(tensor, name, value) - return tensor - - -def _make_sparse_delta_extension( - parameter_name: str, - target: torch.Tensor, - module: object, -) -> Any: - from nemo_rl.models.generation.vllm.vllm_backend import ( - VllmInternalWorkerExtension, - ) - - ext = VllmInternalWorkerExtension.__new__(VllmInternalWorkerExtension) - ext.rank = 1 - ext._direct_sparse_delta_targets = {parameter_name: target} - ext.model_runner = SimpleNamespace( - model=SimpleNamespace(get_submodule=lambda _name: module) - ) - ext._direct_sparse_delta_plan_cache = {} - ext._direct_sparse_delta_verification = None - ext._direct_sparse_delta_verification_candidates = 0 - return ext - - -def _assert_sparse_plan( - ext: Any, - plan: Any, - source_locations: list[int], - expected_locations: list[int], - expected_values: list[float], -) -> None: - assert plan is not None - values = torch.arange(len(source_locations), dtype=torch.float32) - locations, values = ext._local_sparse_delta_update_inputs( - torch.tensor(source_locations), values, plan - ) - assert locations.tolist() == expected_locations - assert values.tolist() == expected_values - - -@pytest.mark.vllm -def test_serialized_sparse_payload_batch_preserves_order(tmp_path) -> None: - from nemo_rl.models.generation.vllm.vllm_backend import ( - VllmInternalWorkerExtension, - ) - - ext = VllmInternalWorkerExtension.__new__(VllmInternalWorkerExtension) - ext.device = torch.device("cpu") - payloads = [ - (torch.tensor([index]), torch.tensor([float(index)]), {"index": index}) - for index in range(3) - ] - paths = [tmp_path / f"{index}.pt" for index in range(3)] - for path, payload in zip(paths, payloads, strict=True): - torch.save(payload, path) - applied: list[Any] = [] - - def apply(payload: Any) -> dict[str, Any]: - applied.append(payload) - return { - "ok": True, - "receiver_sparse_apply_s": 2.0, - } - - ext._apply_sparse_request = apply - result = ext.update_weights_from_sparse_payload_files( - *(str(path) for path in paths) - ) - - assert [item[2]["index"] for item in applied] == [0, 1, 2] - assert all( - torch.equal(item[1], payload[1]) - for item, payload in zip(applied, payloads, strict=True) - ) - assert result["receiver_deserialize_s"] >= 0.0 - assert result["receiver_sparse_apply_s"] == 6.0 - - -@pytest.mark.vllm -def test_direct_sparse_delta_placement() -> None: - qkv_name = "model.layers.0.self_attn.qkv_proj.weight" - qkv_target = _attach_tensor_attrs(torch.zeros(8, 2), output_dim=0) - ext = _make_sparse_delta_extension( - qkv_name, - qkv_target, - SimpleNamespace( - tp_rank=1, - num_kv_head_replicas=2, - _get_shard_offset_mapping=lambda shard: {"q": 0, "k": 4, "v": 6}[shard], - _get_shard_size_mapping=lambda shard: {"q": 4, "k": 2, "v": 2}[shard], - ), - ) - qkv_source = "model.layers.0.self_attn.k_proj.weight" - plan = ext._direct_sparse_delta_qkv_plan( - {"name": qkv_source, "shape": (2, 2)}, qkv_source, {qkv_name: qkv_target} - ) - _assert_sparse_plan(ext, plan, [0, 1, 2, 3], [8, 9, 10, 11], [0.0, 1.0, 2.0, 3.0]) - - merged_name = "model.layers.0.mlp.gate_up_proj.weight" - merged_target = _attach_tensor_attrs(torch.zeros(8, 2), output_dim=0) - ext = _make_sparse_delta_extension( - merged_name, - merged_target, - SimpleNamespace(tp_rank=1, tp_size=2, output_sizes=(8, 8)), - ) - for projection, expected_locations in ( - ("gate", [0, 1, 6, 7]), - ("up", [8, 9, 14, 15]), - ): - source_name = f"model.layers.0.mlp.{projection}_proj.weight" - plan = ext._direct_sparse_delta_target_plan( - {"name": source_name, "shape": (8, 2)}, - {merged_name: merged_target}, - ) - _assert_sparse_plan( - ext, - plan, - [6, 7, 8, 9, 14, 15], - expected_locations, - [2.0, 3.0, 4.0, 5.0], - ) - - expert_name = "model.layers.0.mlp.experts.w13_weight" - expert_target = torch.zeros(2, 4, 2) - expert_module = SimpleNamespace( - tp_rank=1, - moe_config=SimpleNamespace(is_act_and_mul=False), - _map_global_expert_id_to_local_expert_id=lambda expert: ( - 1 if expert == 3 else -1 - ), - ) - ext = _make_sparse_delta_extension( - expert_name, - expert_target, - expert_module, - ) - for projection in ("gate_proj", "up_proj"): - expert_source = f"model.layers.0.mlp.experts.3.{projection}.weight" - plan = ext._direct_sparse_delta_expert_plan( - {"name": expert_source, "shape": (8, 2)}, - expert_source, - {expert_name: expert_target}, - ) - _assert_sparse_plan( - ext, - plan, - [6, 7, 8, 9, 14, 15], - [8, 9, 14, 15], - [2.0, 3.0, 4.0, 5.0], - ) - - w2_target = torch.zeros(2, 2, 4) - ext = _make_sparse_delta_extension(expert_name, w2_target, expert_module) - expert_source = "model.layers.0.mlp.experts.3.down_proj.weight" - plan = ext._direct_sparse_delta_expert_plan( - {"name": expert_source, "shape": (2, 8)}, - expert_source, - {"model.layers.0.mlp.experts.w2_weight": w2_target}, - ) - _assert_sparse_plan(ext, plan, [3, 4, 7, 11, 15], [8, 11, 15], [1.0, 2.0, 4.0]) - - mamba_name = "model.layers.0.mixer.in_proj.weight" - for target_shape, groups, source_locations, expected_locations, values in ( - ((16, 1, 2), 6, [0, 8, 25, 38, 40, 55], [0, 9, 22, 31], [1, 2, 4, 5]), - ( - (14, 2), - 4, - [0, 8, 24, 36, 44, 52, 55], - [0, 8, 16, 20, 24, 27], - [1, 2, 3, 4, 5, 6], - ), - ): - target = _attach_tensor_attrs( - torch.zeros(target_shape), - weight_loader=MethodType(lambda _owner: None, SimpleNamespace()), - ) - ext = _make_sparse_delta_extension( - mamba_name, - target, - SimpleNamespace( - tp_size=2, - intermediate_size=8, - groups_ssm_state_size=groups, - num_heads=4, - ), - ) - plan = ext._direct_sparse_delta_mamba2_plan( - {"name": mamba_name, "shape": (28, 2)}, - mamba_name, - {mamba_name: target}, - ) - _assert_sparse_plan(ext, plan, source_locations, expected_locations, values) - - for attrs, source_shape, source_locations, expected_locations, values in ( - ( - {"output_dim": 0}, - (6, 2), - [0, 1, 6, 7, 10, 11], - [0, 1, 4, 5], - [2, 3, 4, 5], - ), - ( - {"output_dim": 0, "input_dim": 1}, - (3, 4), - [0, 1, 2, 3, 6, 7, 10, 11], - [0, 1, 2, 3, 4, 5], - [2, 3, 4, 5, 6, 7], - ), - ): - target = _attach_tensor_attrs(torch.zeros(3, 2), **attrs, tp_size=2, tp_rank=1) - ext = _make_sparse_delta_extension("down_proj.weight", target, object()) - plan = ext._direct_sparse_delta_shard_plan( - {"name": "down_proj.weight", "shape": source_shape}, target - ) - _assert_sparse_plan(ext, plan, source_locations, expected_locations, values) - - -@pytest.mark.vllm -@pytest.mark.parametrize( - ("initial", "expected_delta", "exact_mismatches", "mismatches"), - [ - (200.0, 4.0, 0, 0), - (2.0, 4.0000005, 1, 0), - (2.0, 5.0, 1, 1), - ], -) -def test_sparse_delta_sample_verification_only_compares_applied_delta( - monkeypatch, - initial: float, - expected_delta: float, - exact_mismatches: int, - mismatches: int, -) -> None: - from nemo_rl.models.generation.vllm.vllm_backend import ( - VllmInternalWorkerExtension, - ) - - monkeypatch.setattr(torch.cuda, "is_available", lambda: False) - monkeypatch.setattr( - "nemo_rl.models.generation.vllm.quantization.fp8.is_fp8_model", - lambda _config: False, - ) - target = torch.tensor([1.0, initial, 3.0, initial]) - ext = VllmInternalWorkerExtension.__new__(VllmInternalWorkerExtension) - ext.model_runner = SimpleNamespace( - model=SimpleNamespace(), - vllm_config=SimpleNamespace( - model_config=SimpleNamespace(architectures=[]), - ), - ) - ext._direct_sparse_delta_targets = {"weight": target} - ext._direct_sparse_delta_plan_cache = { - "weight": ext._make_sparse_delta_target_plan(target, (4,)) - } - ext._direct_sparse_delta_verification = None - ext._direct_sparse_delta_verification_candidates = 0 - payload = encode_sparse_infos( - [("weight", target, torch.tensor([1, 3]), torch.tensor([4.0, 4.0]))], - empty_dtype=target.dtype, - ) - metadata = payload[2] - metadata[0].update( - verification_locations=[1, 3], - verification_deltas=[expected_delta, expected_delta], - ) - - ext._apply_sparse_weight_deltas(payload[:2], metadata) - result = ext.finish_sparse_delta_refit() - - assert torch.equal(target, torch.tensor([1.0, initial + 4.0, 3.0, initial + 4.0])) - assert result["verification_candidates"] == 2 - assert result["verification_samples"] == 2 - assert result["verification_exact_mismatches"] == 2 * exact_mismatches - assert result["verification_mismatches"] == 2 * mismatches - rounded_difference = float((torch.tensor(expected_delta) - 4).abs()) - assert result["verification_max_abs"] == rounded_difference - - @pytest.mark.vllm def test_update_weights_from_collective_processes_weights_after_loading(monkeypatch): from nemo_rl.models.generation.vllm import vllm_backend call_order = [] process_calls = [] - current_configs = [] def process_weights_after_loading(model, model_config, device): - assert current_configs == [ext.model_runner.vllm_config] call_order.append("process") process_calls.append((model, model_config, device)) @@ -426,7 +140,6 @@ def packed_broadcast_consumer(iterator, group, src, post_unpack_func): assert ext.update_weights_from_collective() is True - assert not current_configs assert process_calls == [(ext.model_runner.model, ext.model_config, ext.device)] assert call_order == [ "broadcast", diff --git a/tests/unit/models/generation/test_vllm_generation.py b/tests/unit/models/generation/test_vllm_generation.py index 69add935d26..3017aa9941a 100644 --- a/tests/unit/models/generation/test_vllm_generation.py +++ b/tests/unit/models/generation/test_vllm_generation.py @@ -12,21 +12,14 @@ # See the License for the specific language governing permissions and # limitations under the License. -import asyncio import importlib.util import json import os import sys -import threading -import time import types -from collections.abc import Iterator -from concurrent.futures import Future, ThreadPoolExecutor -from contextlib import contextmanager from copy import deepcopy from pathlib import Path -from typing import Any -from unittest.mock import AsyncMock, MagicMock, call +from unittest.mock import MagicMock import pytest import ray @@ -44,8 +37,6 @@ ) from nemo_rl.models.generation.vllm import VllmConfig, VllmGeneration from nemo_rl.models.generation.vllm.vllm_worker import ( - BaseVllmGenerationWorker, - VllmGenerationWorkerImpl, _resolve_enable_prefix_caching, ) from nemo_rl.models.generation.vllm.vllm_worker_async import ( @@ -54,18 +45,6 @@ ) from nemo_rl.models.policy import LoRAConfig, PolicyConfig from nemo_rl.models.policy.lm_policy import Policy -from nemo_rl.utils.weight_transfer_remote_sparse import ( - G_VLLM_REFIT_API_KEY_HEADER, - G_VLLM_REFIT_FLUSH_PATH, - G_VLLM_REFIT_S3_MANIFEST_PATH, -) -from nemo_rl.utils.weight_transfer_zmq import ( - G_VLLM_REFIT_CHECKSUM_HEADER, - G_VLLM_REFIT_PAYLOAD_HEADER, - G_VLLM_REFIT_PRODUCER_HEADER, - G_VLLM_REFIT_TRANSFER_HEADER, - G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, -) model_name = "Qwen/Qwen3-0.6B" # Define basic vLLM test config @@ -180,485 +159,6 @@ def test_resolve_enable_prefix_caching_uses_cuda_capability_for_auto(monkeypatch assert _resolve_enable_prefix_caching({}) is False -def test_ray_owner_destructors_skip_shutdown_during_interpreter_finalization( - monkeypatch, -): - monkeypatch.setattr(sys, "is_finalizing", lambda: True) - - generation = VllmGeneration.__new__(VllmGeneration) - generation.shutdown = MagicMock() - generation.__del__() - generation.shutdown.assert_not_called() - - policy = Policy.__new__(Policy) - policy.worker_group = MagicMock() - policy.__del__() - policy.worker_group.shutdown.assert_not_called() - - cluster = RayVirtualCluster.__new__(RayVirtualCluster) - cluster.shutdown = MagicMock() - cluster.__del__() - cluster.shutdown.assert_not_called() - - -@contextmanager -def _sparse_refit_worker( - *, batch_size: int = 2, futures: list[Future[dict[str, Any]]] | None = None -) -> Iterator[BaseVllmGenerationWorker]: - worker = BaseVllmGenerationWorker.__new__(BaseVllmGenerationWorker) - worker._refit_apply_queue_condition = threading.Condition() - worker._refit_apply_executor = ThreadPoolExecutor(max_workers=1) - worker._refit_apply_futures = list(futures or []) - worker._refit_apply_pending_payloads = [] - worker._refit_seen_payloads = {} - worker._refit_apply_queue_depth = 2 - worker._refit_apply_batch_size = batch_size - worker.llm = MagicMock() - worker.llm.collective_rpc.return_value = [{"ok": True}] - for future in worker._refit_apply_futures: - future.add_done_callback(worker._notify_refit_apply_waiters) - try: - yield worker - finally: - worker._refit_apply_executor.shutdown(wait=True) - - -def test_sparse_refit_queue_batches_payloads_in_fifo_order() -> None: - applied: list[tuple[bytes, ...]] = [] - - def apply(payloads: tuple[bytes, ...]) -> dict[str, Any]: - applied.append(payloads) - return { - "ok": True, - "payloads": len(payloads), - "receiver_total_s": float(len(payloads)), - } - - with _sparse_refit_worker(batch_size=3) as worker: - worker.update_weights_from_serialized_sparse_payloads = apply - responses = [ - worker._enqueue_sparse_payload_apply( - payload, ("transfer", 0, index), str(index) - ) - for index, payload in enumerate((b"0", b"1", b"2", b"3", b"4")) - ] - response = worker._flush_queued_sparse_payloads() - responses.append(response) - - assert applied == [ - (b"0", b"1", b"2"), - (b"3", b"4"), - ] - assert response["payloads"] == 5 - assert response["batches"] == 2 - assert sum(result.get("receiver_total_s", 0.0) for result in responses) == 5.0 - worker.llm.collective_rpc.assert_called_once_with( - "finish_sparse_delta_refit", args=() - ) - - -def test_sparse_refit_queue_deduplicates_transactional_payloads() -> None: - key = ("transfer", 0, 1) - with _sparse_refit_worker() as worker: - worker.update_weights_from_serialized_sparse_payloads = MagicMock( - return_value={"ok": True, "payloads": 1} - ) - assert worker._enqueue_sparse_payload_apply(b"payload", key, "checksum")["ok"] - duplicate = worker._enqueue_sparse_payload_apply(b"payload", key, "checksum") - assert duplicate == {"ok": True, "payloads": 0, "duplicate": True} - with pytest.raises(ValueError, match="reused with different data"): - worker._enqueue_sparse_payload_apply(b"other", key, "different") - response = worker._flush_queued_sparse_payloads() - - assert response["payloads"] == 1 - assert worker._refit_seen_payloads == {} - - -def test_sparse_refit_queue_does_not_deduplicate_failed_enqueue() -> None: - failed = Future() - failed.set_exception(RuntimeError("prior apply failed")) - with _sparse_refit_worker(futures=[failed]) as worker: - with pytest.raises(RuntimeError, match="prior apply failed"): - worker._enqueue_sparse_payload_apply( - b"payload", ("transfer", 0, 1), "checksum" - ) - - assert worker._refit_seen_payloads == {} - assert worker._refit_apply_pending_payloads == [] - - -def test_sparse_refit_collective_response_merges_verification_metrics() -> None: - response = BaseVllmGenerationWorker._refit_collective_response( - [ - { - "receiver_total_s": 1.0, - "verification_candidates": 4, - "verification_samples": 2, - "verification_exact_mismatches": 1, - "verification_mismatches": 0, - "verification_abs_sum": 0.25, - "verification_max_abs": 0.25, - }, - { - "receiver_total_s": 2.0, - "verification_candidates": 4, - "verification_samples": 3, - "verification_exact_mismatches": 2, - "verification_mismatches": 1, - "verification_abs_sum": 0.5, - "verification_max_abs": 0.4, - }, - ] - ) - - assert response == { - "ok": True, - "receiver_total_s": 2.0, - "verification_candidates": 4, - "verification_samples": 5, - "verification_exact_mismatches": 3, - "verification_mismatches": 1, - "verification_abs_sum": 0.75, - "verification_max_abs": 0.4, - } - - -def test_sparse_refit_queue_releases_condition_while_backpressured() -> None: - first: Future[dict[str, Any]] = Future() - second: Future[dict[str, Any]] = Future() - started = threading.Event() - - with _sparse_refit_worker(futures=[first, second]) as worker: - with ThreadPoolExecutor(max_workers=1) as callers: - call = callers.submit( - lambda: ( - started.set(), - worker._enqueue_sparse_payload_apply( - b"payload", ("transfer", 0, 1), "checksum" - ), - )[1] - ) - assert started.wait(timeout=1.0) - time.sleep(0.05) - assert worker._refit_apply_queue_condition.acquire(timeout=1.0) - worker._refit_apply_queue_condition.release() - first.set_result({"ok": True, "payloads": 1}) - assert call.result(timeout=1.0)["ok"] - - -def test_sparse_refit_batch_uses_one_collective_rpc(tmp_path: Path) -> None: - worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) - staged_payloads: list[bytes] = [] - - def collective_rpc(method, args): - assert method == "update_weights_from_sparse_payload_files" - staged_payloads.extend(Path(path).read_bytes() for path in args) - return [{"ok": True, "receiver_total_s": 1.0}] - - worker.llm = MagicMock(collective_rpc=MagicMock(side_effect=collective_rpc)) - worker._refit_workers_share_node = True - worker._refit_batch_staging_dir = str(tmp_path) - payloads = (b"0", b"1", b"2") - - response = worker.update_weights_from_serialized_sparse_payloads(payloads) - - assert staged_payloads == list(payloads) - assert not list(tmp_path.iterdir()) - worker.llm.collective_rpc.assert_called_once() - assert response == {"ok": True, "receiver_total_s": 1.0, "payloads": 3} - - -def test_sparse_refit_batch_drains_workers_before_error_cleanup(tmp_path: Path) -> None: - worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) - staged_paths: tuple[str, ...] = () - - def collective_rpc(method, args): - nonlocal staged_paths - if method == "update_weights_from_sparse_payload_files": - staged_paths = args - raise RuntimeError("apply failed") - assert method == "synchronize_device" - assert all(Path(path).is_file() for path in staged_paths) - return [True] - - worker.llm = MagicMock(collective_rpc=MagicMock(side_effect=collective_rpc)) - worker._refit_workers_share_node = True - worker._refit_batch_staging_dir = str(tmp_path) - - with pytest.raises(RuntimeError, match="apply failed"): - worker.update_weights_from_serialized_sparse_payloads((b"0", b"1")) - - assert [call.args[0] for call in worker.llm.collective_rpc.call_args_list] == [ - "update_weights_from_sparse_payload_files", - "synchronize_device", - ] - assert not list(tmp_path.iterdir()) - - -def test_sparse_refit_batch_falls_back_across_nodes() -> None: - worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) - worker._refit_workers_share_node = False - worker.llm = MagicMock( - collective_rpc=MagicMock(return_value=[{"ok": True, "receiver_total_s": 1.0}]) - ) - - response = worker.update_weights_from_serialized_sparse_payloads((b"0", b"1", b"2")) - - assert worker.llm.collective_rpc.call_args_list == [ - call("update_weights_from_serialized_sparse_payload", args=(payload,)) - for payload in (b"0", b"1", b"2") - ] - assert response == {"ok": True, "receiver_total_s": 3.0, "payloads": 3} - - -@pytest.mark.asyncio -async def test_async_sparse_refit_batch_bridges_to_async_collective( - tmp_path: Path, -) -> None: - worker = VllmAsyncGenerationWorkerImpl.__new__(VllmAsyncGenerationWorkerImpl) - staged_payloads: list[bytes] = [] - - class AsyncLlm: - async def collective_rpc( - self, method: str, args: tuple[Any, ...] - ) -> list[dict[str, Any]]: - assert method == "update_weights_from_sparse_payload_files" - staged_payloads.extend(Path(path).read_bytes() for path in args) - return [{"ok": True, "receiver_total_s": 1.0}] - - worker.llm = AsyncLlm() - worker._refit_async_loop = asyncio.get_running_loop() - worker._refit_workers_share_node = True - worker._refit_batch_staging_dir = str(tmp_path) - - response = await asyncio.to_thread( - worker.update_weights_from_serialized_sparse_payloads, - (b"0", b"1"), - ) - - assert staged_payloads == [b"0", b"1"] - assert response == {"ok": True, "receiver_total_s": 1.0, "payloads": 2} - - -@pytest.mark.asyncio -async def test_sparse_refit_payload_handlers_decode_and_enqueue(monkeypatch) -> None: - worker = BaseVllmGenerationWorker.__new__(BaseVllmGenerationWorker) - enqueue = MagicMock( - side_effect=[ - {"ok": True, "payloads": 1}, - {"ok": True, "payloads": 1}, - ] - ) - worker._enqueue_sparse_payload_apply = enqueue - monkeypatch.setattr( - "nemo_rl.models.generation.vllm.vllm_worker.download_s3_refit_payload", - lambda _manifest: b"s3-payload", - ) - monkeypatch.setattr( - "nemo_rl.models.generation.vllm.vllm_worker.decode_sparse_payload", - lambda _body, _checksum: b"zmq-payload", - ) - - s3_result = await worker._apply_s3_manifest_payload( - {"key": "object-key", "checksum": "checksum"} - ) - assert s3_result["ok"] - assert s3_result["receiver_s3_download_s"] >= 0.0 - - invalid_request = types.SimpleNamespace(headers={}, body=AsyncMock()) - with pytest.raises(ValueError, match="Missing or invalid"): - await worker._apply_zmq_payload(invalid_request) - - request = types.SimpleNamespace( - headers={ - G_VLLM_REFIT_TRANSFER_HEADER: "transfer", - G_VLLM_REFIT_PRODUCER_HEADER: "2", - G_VLLM_REFIT_PAYLOAD_HEADER: "3", - G_VLLM_REFIT_CHECKSUM_HEADER: "checksum", - }, - body=AsyncMock(return_value=b"compressed"), - ) - zmq_result = await worker._apply_zmq_payload(request) - assert zmq_result["ok"] - assert zmq_result["receiver_zmq_decode_s"] >= 0.0 - assert enqueue.call_args_list == [ - call(b"s3-payload", ("object-key", -1, -1), "checksum"), - call(b"zmq-payload", ("transfer", 2, 3), "checksum"), - ] - - -def test_sparse_refit_api_auth_dispatch_and_error_mapping(monkeypatch) -> None: - from fastapi import FastAPI - from fastapi.testclient import TestClient - - monkeypatch.setenv("NRL_TEST_REFIT_KEY", "secret") - worker = BaseVllmGenerationWorker.__new__(BaseVllmGenerationWorker) - worker.cfg = { - "vllm_cfg": { - "async_engine": True, - "http_refit_api_key_env_var": "NRL_TEST_REFIT_KEY", - } - } - worker._refit_async_loop = None - worker._apply_s3_manifest_payload = AsyncMock( - return_value={"ok": True, "payloads": 1} - ) - worker._apply_zmq_payload = AsyncMock(side_effect=RuntimeError("apply failed")) - worker._flush_queued_sparse_payloads = MagicMock( - return_value={"ok": True, "payloads": 2} - ) - app = FastAPI() - worker._setup_vllm_refit_api_server(app) - headers = {G_VLLM_REFIT_API_KEY_HEADER: "secret"} - - with TestClient(app) as client: - unauthorized = client.post(G_VLLM_REFIT_S3_MANIFEST_PATH, json={}) - s3_response = client.post( - G_VLLM_REFIT_S3_MANIFEST_PATH, - json={"key": "key"}, - headers=headers, - ) - flush_response = client.post(G_VLLM_REFIT_FLUSH_PATH, headers=headers) - zmq_response = client.post( - G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, - content=b"payload", - headers=headers, - ) - - assert unauthorized.status_code == 403 - assert unauthorized.json() == {"ok": False, "error": "unauthorized"} - assert s3_response.status_code == 200 - assert s3_response.json() == {"ok": True, "payloads": 1} - assert flush_response.status_code == 200 - assert flush_response.json() == {"ok": True, "payloads": 2} - assert zmq_response.status_code == 500 - assert zmq_response.json() == {"ok": False, "error": "apply failed"} - assert worker._refit_async_loop is not None - worker._apply_s3_manifest_payload.assert_awaited_once_with({"key": "key"}) - worker._apply_zmq_payload.assert_awaited_once() - worker._flush_queued_sparse_payloads.assert_called_once_with() - - -def test_sync_sparse_refit_server_shutdown_cleans_transport_resources( - monkeypatch, -) -> None: - import uvicorn - - from nemo_rl.models.generation.vllm import vllm_worker as worker_module - - configs = [] - servers = [] - - def make_config(app, **kwargs): - config = types.SimpleNamespace(app=app, **kwargs) - configs.append(config) - return config - - class Server: - def __init__(self, config) -> None: - self.config = config - self.should_exit = False - self.ran = threading.Event() - servers.append(self) - - def run(self) -> None: - self.ran.set() - - monkeypatch.setattr(uvicorn, "Config", make_config) - monkeypatch.setattr(uvicorn, "Server", Server) - monkeypatch.setattr(worker_module, "_get_free_port_local", lambda *_args: 12345) - monkeypatch.setattr(worker_module, "_get_node_ip_local", lambda: "10.0.0.1") - collect = MagicMock() - empty_cache = MagicMock() - monkeypatch.setattr(worker_module.gc, "collect", collect) - monkeypatch.setattr(worker_module.torch.cuda, "empty_cache", empty_cache) - - worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) - worker.cfg = { - "vllm_cfg": {"http_refit_server_port": None}, - "port_range_low": 10000, - "port_range_high": 11000, - } - worker._setup_vllm_refit_api_server = MagicMock() - worker._refit_http_server = None - worker._setup_vllm_refit_server() - - assert len(configs) == 1 - assert configs[0].host == "0.0.0.0" - assert configs[0].port == 12345 - assert servers[0].ran.wait(timeout=1.0) - assert worker.report_refit_server_base_url() == "http://10.0.0.1:12345" - worker._setup_vllm_refit_api_server.assert_called_once_with(configs[0].app) - - relay = MagicMock() - llm = MagicMock() - worker._zmq_refit_server = (relay, "tcp://relay") - worker._flush_queued_sparse_payloads = MagicMock() - worker._refit_apply_executor = MagicMock() - worker.llm = llm - worker.tokenizer = object() - - assert worker.shutdown() is True - relay.close.assert_called_once_with() - worker._flush_queued_sparse_payloads.assert_called_once_with() - worker._refit_apply_executor.shutdown.assert_called_once_with(wait=True) - assert servers[0].should_exit is True - assert worker._refit_http_server is None - llm.collective_rpc.assert_called_once_with("cleanup", args=()) - assert worker.llm is None - assert worker.tokenizer is None - collect.assert_called_once_with() - empty_cache.assert_called_once_with() - - -@pytest.mark.asyncio -async def test_async_sparse_refit_post_init_records_worker_locality() -> None: - worker = VllmAsyncGenerationWorkerImpl.__new__(VllmAsyncGenerationWorkerImpl) - worker.cfg = {"refit_transport": "vllm_zmq_sparse"} - worker._mtp_load_from_disk = False - worker.report_device_id_async = AsyncMock(return_value=["0"]) - worker.llm = MagicMock() - worker.llm.collective_rpc = AsyncMock(return_value=["node-0", "node-0"]) - - await worker.post_init_async() - - assert worker.vllm_device_ids == ["0"] - assert worker._refit_workers_share_node is True - assert worker.llm.collective_rpc.await_args_list == [ - call("bind_numa", args=()), - call("report_node_hostname", args=()), - ] - - -def test_async_sparse_refit_exposes_zmq_relay(monkeypatch) -> None: - server = MagicMock() - server_type = MagicMock(return_value=server) - monkeypatch.setattr( - "nemo_rl.models.generation.vllm.vllm_worker.ZmqSparseRefitServer", - server_type, - ) - monkeypatch.setattr( - "nemo_rl.models.generation.vllm.vllm_worker._get_free_port_local", - lambda *_args: 12345, - ) - monkeypatch.setattr( - "nemo_rl.models.generation.vllm.vllm_worker._get_node_ip_local", - lambda: "10.0.0.1", - ) - worker = VllmAsyncGenerationWorkerImpl.__new__(VllmAsyncGenerationWorkerImpl) - worker.cfg = {"vllm_cfg": {"zmq_refit_server_port": None}} - worker._zmq_refit_server = None - - assert worker.start_zmq_sparse_refit_relay(["http://receiver"]) == ( - "tcp://10.0.0.1:12345" - ) - server.start.assert_called_once_with() - - worker.stop_zmq_sparse_refit_relay() - server.close.assert_called_once_with() - assert worker._zmq_refit_server is None - - basic_lora_test_config: LoRAConfig = { "enabled": False, "target_modules": [], diff --git a/tests/unit/models/generation/test_vllm_sparse_delta.py b/tests/unit/models/generation/test_vllm_sparse_delta.py new file mode 100644 index 00000000000..f23202ea822 --- /dev/null +++ b/tests/unit/models/generation/test_vllm_sparse_delta.py @@ -0,0 +1,288 @@ +# 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 MethodType, SimpleNamespace +from typing import Any + +import pytest +import torch + +from nemo_rl.models.generation.vllm.vllm_sparse_delta import VllmSparseDeltaApplier +from nemo_rl.utils.weight_transfer_sparse_codec import encode_sparse_infos + + +def _attach_tensor_attrs(tensor: torch.Tensor, **attrs: object) -> torch.Tensor: + for name, value in attrs.items(): + setattr(tensor, name, value) + return tensor + + +def _make_sparse_delta_extension( + parameter_name: str, + target: torch.Tensor, + module: object, +) -> Any: + model_runner = SimpleNamespace( + model=SimpleNamespace(get_submodule=lambda _name: module) + ) + ext = VllmSparseDeltaApplier( + model_runner, + torch.device("cpu"), + rank=1, + ) + ext._direct_sparse_delta_targets = {parameter_name: target} + return ext + + +def _assert_sparse_plan( + ext: Any, + plan: Any, + source_locations: list[int], + expected_locations: list[int], + expected_values: list[float], +) -> None: + assert plan is not None + values = torch.arange(len(source_locations), dtype=torch.float32) + locations, values = ext._local_sparse_delta_update_inputs( + torch.tensor(source_locations), values, plan + ) + assert locations.tolist() == expected_locations + assert values.tolist() == expected_values + + +@pytest.mark.vllm +def test_serialized_sparse_payload_batch_preserves_order(tmp_path) -> None: + ext = VllmSparseDeltaApplier(SimpleNamespace(), torch.device("cpu")) + payloads = [ + (torch.tensor([index]), torch.tensor([float(index)]), {"index": index}) + for index in range(3) + ] + paths = [tmp_path / f"{index}.pt" for index in range(3)] + for path, payload in zip(paths, payloads, strict=True): + torch.save(payload, path) + applied: list[Any] = [] + + def apply(payload: Any) -> dict[str, Any]: + applied.append(payload) + return { + "ok": True, + "receiver_sparse_apply_s": 2.0, + } + + ext._apply_sparse_request = apply + result = ext.update_weights_from_sparse_payload_files( + *(str(path) for path in paths) + ) + + assert [item[2]["index"] for item in applied] == [0, 1, 2] + assert all( + torch.equal(item[1], payload[1]) + for item, payload in zip(applied, payloads, strict=True) + ) + assert result["receiver_deserialize_s"] >= 0.0 + assert result["receiver_sparse_apply_s"] == 6.0 + + +@pytest.mark.vllm +def test_direct_sparse_delta_placement() -> None: + qkv_name = "model.layers.0.self_attn.qkv_proj.weight" + qkv_target = _attach_tensor_attrs(torch.zeros(8, 2), output_dim=0) + ext = _make_sparse_delta_extension( + qkv_name, + qkv_target, + SimpleNamespace( + tp_rank=1, + num_kv_head_replicas=2, + _get_shard_offset_mapping=lambda shard: {"q": 0, "k": 4, "v": 6}[shard], + _get_shard_size_mapping=lambda shard: {"q": 4, "k": 2, "v": 2}[shard], + ), + ) + qkv_source = "model.layers.0.self_attn.k_proj.weight" + plan = ext._direct_sparse_delta_qkv_plan( + {"name": qkv_source, "shape": (2, 2)}, qkv_source, {qkv_name: qkv_target} + ) + _assert_sparse_plan(ext, plan, [0, 1, 2, 3], [8, 9, 10, 11], [0.0, 1.0, 2.0, 3.0]) + + merged_name = "model.layers.0.mlp.gate_up_proj.weight" + merged_target = _attach_tensor_attrs(torch.zeros(8, 2), output_dim=0) + ext = _make_sparse_delta_extension( + merged_name, + merged_target, + SimpleNamespace(tp_rank=1, tp_size=2, output_sizes=(8, 8)), + ) + for projection, expected_locations in ( + ("gate", [0, 1, 6, 7]), + ("up", [8, 9, 14, 15]), + ): + source_name = f"model.layers.0.mlp.{projection}_proj.weight" + plan = ext._direct_sparse_delta_target_plan( + {"name": source_name, "shape": (8, 2)}, + {merged_name: merged_target}, + ) + _assert_sparse_plan( + ext, + plan, + [6, 7, 8, 9, 14, 15], + expected_locations, + [2.0, 3.0, 4.0, 5.0], + ) + + expert_name = "model.layers.0.mlp.experts.w13_weight" + expert_target = torch.zeros(2, 4, 2) + expert_module = SimpleNamespace( + tp_rank=1, + moe_config=SimpleNamespace(is_act_and_mul=False), + _map_global_expert_id_to_local_expert_id=lambda expert: ( + 1 if expert == 3 else -1 + ), + ) + ext = _make_sparse_delta_extension( + expert_name, + expert_target, + expert_module, + ) + for projection in ("gate_proj", "up_proj"): + expert_source = f"model.layers.0.mlp.experts.3.{projection}.weight" + plan = ext._direct_sparse_delta_expert_plan( + {"name": expert_source, "shape": (8, 2)}, + expert_source, + {expert_name: expert_target}, + ) + _assert_sparse_plan( + ext, + plan, + [6, 7, 8, 9, 14, 15], + [8, 9, 14, 15], + [2.0, 3.0, 4.0, 5.0], + ) + + w2_target = torch.zeros(2, 2, 4) + ext = _make_sparse_delta_extension(expert_name, w2_target, expert_module) + expert_source = "model.layers.0.mlp.experts.3.down_proj.weight" + plan = ext._direct_sparse_delta_expert_plan( + {"name": expert_source, "shape": (2, 8)}, + expert_source, + {"model.layers.0.mlp.experts.w2_weight": w2_target}, + ) + _assert_sparse_plan(ext, plan, [3, 4, 7, 11, 15], [8, 11, 15], [1.0, 2.0, 4.0]) + + mamba_name = "model.layers.0.mixer.in_proj.weight" + for target_shape, groups, source_locations, expected_locations, values in ( + ((16, 1, 2), 6, [0, 8, 25, 38, 40, 55], [0, 9, 22, 31], [1, 2, 4, 5]), + ( + (14, 2), + 4, + [0, 8, 24, 36, 44, 52, 55], + [0, 8, 16, 20, 24, 27], + [1, 2, 3, 4, 5, 6], + ), + ): + target = _attach_tensor_attrs( + torch.zeros(target_shape), + weight_loader=MethodType(lambda _owner: None, SimpleNamespace()), + ) + ext = _make_sparse_delta_extension( + mamba_name, + target, + SimpleNamespace( + tp_size=2, + intermediate_size=8, + groups_ssm_state_size=groups, + num_heads=4, + ), + ) + plan = ext._direct_sparse_delta_mamba2_plan( + {"name": mamba_name, "shape": (28, 2)}, + mamba_name, + {mamba_name: target}, + ) + _assert_sparse_plan(ext, plan, source_locations, expected_locations, values) + + for attrs, source_shape, source_locations, expected_locations, values in ( + ( + {"output_dim": 0}, + (6, 2), + [0, 1, 6, 7, 10, 11], + [0, 1, 4, 5], + [2, 3, 4, 5], + ), + ( + {"output_dim": 0, "input_dim": 1}, + (3, 4), + [0, 1, 2, 3, 6, 7, 10, 11], + [0, 1, 2, 3, 4, 5], + [2, 3, 4, 5, 6, 7], + ), + ): + target = _attach_tensor_attrs(torch.zeros(3, 2), **attrs, tp_size=2, tp_rank=1) + ext = _make_sparse_delta_extension("down_proj.weight", target, object()) + plan = ext._direct_sparse_delta_shard_plan( + {"name": "down_proj.weight", "shape": source_shape}, target + ) + _assert_sparse_plan(ext, plan, source_locations, expected_locations, values) + + +@pytest.mark.vllm +@pytest.mark.parametrize( + ("initial", "expected_delta", "exact_mismatches", "mismatches"), + [ + (200.0, 4.0, 0, 0), + (2.0, 4.0000005, 1, 0), + (2.0, 5.0, 1, 1), + ], +) +def test_sparse_delta_sample_verification_only_compares_applied_delta( + monkeypatch, + initial: float, + expected_delta: float, + exact_mismatches: int, + mismatches: int, +) -> None: + monkeypatch.setattr(torch.cuda, "is_available", lambda: False) + monkeypatch.setattr( + "nemo_rl.models.generation.vllm.quantization.fp8.is_fp8_model", + lambda _config: False, + ) + target = torch.tensor([1.0, initial, 3.0, initial]) + model_runner = SimpleNamespace( + model=SimpleNamespace(), + vllm_config=SimpleNamespace( + model_config=SimpleNamespace(architectures=[]), + ), + ) + ext = VllmSparseDeltaApplier(model_runner, torch.device("cpu")) + ext._direct_sparse_delta_targets = {"weight": target} + ext._direct_sparse_delta_plan_cache = { + "weight": ext._make_sparse_delta_target_plan(target, (4,)) + } + payload = encode_sparse_infos( + [("weight", target, torch.tensor([1, 3]), torch.tensor([4.0, 4.0]))], + empty_dtype=target.dtype, + ) + metadata = payload[2] + metadata[0].update( + verification_locations=[1, 3], + verification_deltas=[expected_delta, expected_delta], + ) + + ext._apply_sparse_weight_deltas(payload[:2], metadata) + result = ext.finish_sparse_delta_refit() + + assert torch.equal(target, torch.tensor([1.0, initial + 4.0, 3.0, initial + 4.0])) + assert result["verification_candidates"] == 2 + assert result["verification_samples"] == 2 + assert result["verification_exact_mismatches"] == 2 * exact_mismatches + assert result["verification_mismatches"] == 2 * mismatches + rounded_difference = float((torch.tensor(expected_delta) - 4).abs()) + assert result["verification_max_abs"] == rounded_difference diff --git a/tests/unit/models/generation/test_vllm_sparse_refit.py b/tests/unit/models/generation/test_vllm_sparse_refit.py new file mode 100644 index 00000000000..21a821643d8 --- /dev/null +++ b/tests/unit/models/generation/test_vllm_sparse_refit.py @@ -0,0 +1,491 @@ +# 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 asyncio +import threading +import time +from collections.abc import Iterator +from concurrent.futures import Future, ThreadPoolExecutor +from contextlib import contextmanager +from pathlib import Path +from types import SimpleNamespace +from typing import Any +from unittest.mock import AsyncMock, MagicMock, call + +import pytest + +from nemo_rl.models.generation.vllm.vllm_sparse_refit import ( + VllmSparseRefitReceiver, +) +from nemo_rl.models.generation.vllm.vllm_worker_async import ( + VllmAsyncGenerationWorkerImpl, +) +from nemo_rl.utils.weight_transfer_remote_sparse import ( + G_VLLM_REFIT_API_KEY_HEADER, + G_VLLM_REFIT_FLUSH_PATH, + G_VLLM_REFIT_S3_MANIFEST_PATH, +) +from nemo_rl.utils.weight_transfer_zmq import ( + G_VLLM_REFIT_CHECKSUM_HEADER, + G_VLLM_REFIT_PAYLOAD_HEADER, + G_VLLM_REFIT_PRODUCER_HEADER, + G_VLLM_REFIT_TRANSFER_HEADER, + G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, +) + + +@contextmanager +def _sparse_refit_receiver( + *, + batch_size: int = 2, + futures: list[Future[dict[str, Any]]] | None = None, + async_engine: bool = False, + config: dict[str, Any] | None = None, +) -> Iterator[VllmSparseRefitReceiver]: + owner = SimpleNamespace( + cfg=config or {"vllm_cfg": {"async_engine": async_engine}}, + llm=MagicMock(), + ) + owner.llm.collective_rpc.return_value = [{"ok": True}] + receiver = VllmSparseRefitReceiver(owner) + receiver._refit_apply_batch_size = batch_size + receiver._refit_apply_futures = list(futures or []) + for future in receiver._refit_apply_futures: + future.add_done_callback(receiver._notify_refit_apply_waiters) + try: + yield receiver + finally: + receiver._refit_apply_executor.shutdown(wait=True) + + +def test_sparse_refit_queue_batches_payloads_in_fifo_order() -> None: + applied: list[tuple[bytes, ...]] = [] + + def apply(payloads: tuple[bytes, ...]) -> dict[str, Any]: + applied.append(payloads) + return { + "ok": True, + "payloads": len(payloads), + "receiver_total_s": float(len(payloads)), + } + + with _sparse_refit_receiver(batch_size=3) as receiver: + receiver.update_weights_from_serialized_sparse_payloads = apply + responses = [ + receiver._enqueue_sparse_payload_apply( + payload, ("transfer", 0, index), str(index) + ) + for index, payload in enumerate((b"0", b"1", b"2", b"3", b"4")) + ] + response = receiver._flush_queued_sparse_payloads() + responses.append(response) + + assert applied == [(b"0", b"1", b"2"), (b"3", b"4")] + assert response["payloads"] == 5 + assert response["batches"] == 2 + assert sum(result.get("receiver_total_s", 0.0) for result in responses) == 5.0 + receiver.llm.collective_rpc.assert_called_once_with( + "finish_sparse_delta_refit", args=() + ) + + +def test_sparse_refit_queue_deduplicates_transactional_payloads() -> None: + key = ("transfer", 0, 1) + with _sparse_refit_receiver() as receiver: + receiver.update_weights_from_serialized_sparse_payloads = MagicMock( + return_value={"ok": True, "payloads": 1} + ) + assert receiver._enqueue_sparse_payload_apply(b"payload", key, "checksum")["ok"] + duplicate = receiver._enqueue_sparse_payload_apply(b"payload", key, "checksum") + assert duplicate == {"ok": True, "payloads": 0, "duplicate": True} + with pytest.raises(ValueError, match="reused with different data"): + receiver._enqueue_sparse_payload_apply(b"other", key, "different") + response = receiver._flush_queued_sparse_payloads() + + assert response["payloads"] == 1 + assert receiver._refit_seen_payloads == {} + + +def test_sparse_refit_queue_does_not_deduplicate_failed_enqueue() -> None: + failed = Future() + failed.set_exception(RuntimeError("prior apply failed")) + with _sparse_refit_receiver(futures=[failed]) as receiver: + with pytest.raises(RuntimeError, match="prior apply failed"): + receiver._enqueue_sparse_payload_apply( + b"payload", ("transfer", 0, 1), "checksum" + ) + + assert receiver._refit_seen_payloads == {} + assert receiver._refit_apply_pending_payloads == [] + + +def test_sparse_refit_collective_response_merges_verification_metrics() -> None: + response = VllmSparseRefitReceiver._refit_collective_response( + [ + { + "receiver_total_s": 1.0, + "verification_candidates": 4, + "verification_samples": 2, + "verification_exact_mismatches": 1, + "verification_mismatches": 0, + "verification_abs_sum": 0.25, + "verification_max_abs": 0.25, + }, + { + "receiver_total_s": 2.0, + "verification_candidates": 4, + "verification_samples": 3, + "verification_exact_mismatches": 2, + "verification_mismatches": 1, + "verification_abs_sum": 0.5, + "verification_max_abs": 0.4, + }, + ] + ) + + assert response == { + "ok": True, + "receiver_total_s": 2.0, + "verification_candidates": 4, + "verification_samples": 5, + "verification_exact_mismatches": 3, + "verification_mismatches": 1, + "verification_abs_sum": 0.75, + "verification_max_abs": 0.4, + } + + +def test_sparse_refit_queue_releases_condition_while_backpressured() -> None: + first: Future[dict[str, Any]] = Future() + second: Future[dict[str, Any]] = Future() + started = threading.Event() + + with _sparse_refit_receiver(futures=[first, second]) as receiver: + with ThreadPoolExecutor(max_workers=1) as callers: + pending_call = callers.submit( + lambda: ( + started.set(), + receiver._enqueue_sparse_payload_apply( + b"payload", ("transfer", 0, 1), "checksum" + ), + )[1] + ) + assert started.wait(timeout=1.0) + time.sleep(0.05) + assert receiver._refit_apply_queue_condition.acquire(timeout=1.0) + receiver._refit_apply_queue_condition.release() + first.set_result({"ok": True, "payloads": 1}) + assert pending_call.result(timeout=1.0)["ok"] + + +def test_sparse_refit_batch_uses_one_collective_rpc(tmp_path: Path) -> None: + with _sparse_refit_receiver() as receiver: + staged_payloads: list[bytes] = [] + + def collective_rpc(method: str, args: tuple[str, ...]) -> list[dict[str, Any]]: + assert method == "update_weights_from_sparse_payload_files" + staged_payloads.extend(Path(path).read_bytes() for path in args) + return [{"ok": True, "receiver_total_s": 1.0}] + + receiver._worker.llm = MagicMock( + collective_rpc=MagicMock(side_effect=collective_rpc) + ) + receiver._refit_workers_share_node = True + receiver._refit_batch_staging_dir = str(tmp_path) + response = receiver.update_weights_from_serialized_sparse_payloads( + (b"0", b"1", b"2") + ) + + assert staged_payloads == [b"0", b"1", b"2"] + assert not list(tmp_path.iterdir()) + receiver.llm.collective_rpc.assert_called_once() + assert response == {"ok": True, "receiver_total_s": 1.0, "payloads": 3} + + +def test_sparse_refit_batch_drains_workers_before_error_cleanup(tmp_path: Path) -> None: + with _sparse_refit_receiver() as receiver: + staged_paths: tuple[str, ...] = () + + def collective_rpc(method: str, args: tuple[str, ...]) -> list[Any]: + nonlocal staged_paths + if method == "update_weights_from_sparse_payload_files": + staged_paths = args + raise RuntimeError("apply failed") + assert method == "synchronize_device" + assert all(Path(path).is_file() for path in staged_paths) + return [True] + + receiver._worker.llm = MagicMock( + collective_rpc=MagicMock(side_effect=collective_rpc) + ) + receiver._refit_workers_share_node = True + receiver._refit_batch_staging_dir = str(tmp_path) + + with pytest.raises(RuntimeError, match="apply failed"): + receiver.update_weights_from_serialized_sparse_payloads((b"0", b"1")) + + assert [ + entry.args[0] for entry in receiver.llm.collective_rpc.call_args_list + ] == [ + "update_weights_from_sparse_payload_files", + "synchronize_device", + ] + assert not list(tmp_path.iterdir()) + + +def test_sparse_refit_batch_falls_back_across_nodes() -> None: + with _sparse_refit_receiver() as receiver: + receiver._refit_workers_share_node = False + receiver._worker.llm = MagicMock( + collective_rpc=MagicMock( + return_value=[{"ok": True, "receiver_total_s": 1.0}] + ) + ) + + response = receiver.update_weights_from_serialized_sparse_payloads( + (b"0", b"1", b"2") + ) + + assert receiver.llm.collective_rpc.call_args_list == [ + call("update_weights_from_serialized_sparse_payload", args=(payload,)) + for payload in (b"0", b"1", b"2") + ] + assert response == {"ok": True, "receiver_total_s": 3.0, "payloads": 3} + + +@pytest.mark.asyncio +async def test_async_sparse_refit_batch_bridges_to_async_collective( + tmp_path: Path, +) -> None: + with _sparse_refit_receiver(async_engine=True) as receiver: + staged_payloads: list[bytes] = [] + + class AsyncLlm: + async def collective_rpc( + self, method: str, args: tuple[str, ...] + ) -> list[dict[str, Any]]: + assert method == "update_weights_from_sparse_payload_files" + staged_payloads.extend(Path(path).read_bytes() for path in args) + return [{"ok": True, "receiver_total_s": 1.0}] + + receiver._worker.llm = AsyncLlm() + receiver._refit_async_loop = asyncio.get_running_loop() + receiver._refit_workers_share_node = True + receiver._refit_batch_staging_dir = str(tmp_path) + + response = await asyncio.to_thread( + receiver.update_weights_from_serialized_sparse_payloads, + (b"0", b"1"), + ) + + assert staged_payloads == [b"0", b"1"] + assert response == {"ok": True, "receiver_total_s": 1.0, "payloads": 2} + + +@pytest.mark.asyncio +async def test_sparse_refit_payload_handlers_decode_and_enqueue(monkeypatch) -> None: + with _sparse_refit_receiver() as receiver: + enqueue = MagicMock( + side_effect=[ + {"ok": True, "payloads": 1}, + {"ok": True, "payloads": 1}, + ] + ) + receiver._enqueue_sparse_payload_apply = enqueue + monkeypatch.setattr( + "nemo_rl.models.generation.vllm.vllm_sparse_refit.download_s3_refit_payload", + lambda _manifest: b"s3-payload", + ) + monkeypatch.setattr( + "nemo_rl.models.generation.vllm.vllm_sparse_refit.decode_sparse_payload", + lambda _body, _checksum: b"zmq-payload", + ) + + s3_result = await receiver._apply_s3_manifest_payload( + {"key": "object-key", "checksum": "checksum"} + ) + assert s3_result["ok"] + + invalid_request = SimpleNamespace(headers={}, body=AsyncMock()) + with pytest.raises(ValueError, match="Missing or invalid"): + await receiver._apply_zmq_payload(invalid_request) + + request = SimpleNamespace( + headers={ + G_VLLM_REFIT_TRANSFER_HEADER: "transfer", + G_VLLM_REFIT_PRODUCER_HEADER: "2", + G_VLLM_REFIT_PAYLOAD_HEADER: "3", + G_VLLM_REFIT_CHECKSUM_HEADER: "checksum", + }, + body=AsyncMock(return_value=b"compressed"), + ) + zmq_result = await receiver._apply_zmq_payload(request) + assert zmq_result["ok"] + assert enqueue.call_args_list == [ + call(b"s3-payload", ("object-key", -1, -1), "checksum"), + call(b"zmq-payload", ("transfer", 2, 3), "checksum"), + ] + + +def test_sparse_refit_api_auth_dispatch_and_error_mapping(monkeypatch) -> None: + from fastapi import FastAPI + from fastapi.testclient import TestClient + + monkeypatch.setenv("NRL_TEST_REFIT_KEY", "secret") + config = { + "vllm_cfg": { + "async_engine": True, + "http_refit_api_key_env_var": "NRL_TEST_REFIT_KEY", + } + } + with _sparse_refit_receiver(async_engine=True, config=config) as receiver: + receiver._apply_s3_manifest_payload = AsyncMock( + return_value={"ok": True, "payloads": 1} + ) + receiver._apply_zmq_payload = AsyncMock( + side_effect=RuntimeError("apply failed") + ) + receiver._flush_queued_sparse_payloads = MagicMock( + return_value={"ok": True, "payloads": 2} + ) + app = FastAPI() + receiver.setup_api_server(app) + headers = {G_VLLM_REFIT_API_KEY_HEADER: "secret"} + + with TestClient(app) as client: + unauthorized = client.post(G_VLLM_REFIT_S3_MANIFEST_PATH, json={}) + s3_response = client.post( + G_VLLM_REFIT_S3_MANIFEST_PATH, + json={"key": "key"}, + headers=headers, + ) + flush_response = client.post(G_VLLM_REFIT_FLUSH_PATH, headers=headers) + zmq_response = client.post( + G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, + content=b"payload", + headers=headers, + ) + + assert unauthorized.status_code == 403 + assert s3_response.status_code == 200 + assert flush_response.status_code == 200 + assert zmq_response.status_code == 500 + assert receiver._refit_async_loop is not None + receiver._apply_s3_manifest_payload.assert_awaited_once_with({"key": "key"}) + receiver._apply_zmq_payload.assert_awaited_once() + receiver._flush_queued_sparse_payloads.assert_called_once_with() + + +def test_sync_sparse_refit_server_shutdown_cleans_transport_resources( + monkeypatch, +) -> None: + import uvicorn + + from nemo_rl.models.generation.vllm import vllm_sparse_refit as refit_module + + configs: list[Any] = [] + servers: list[Any] = [] + + def make_config(app: Any, **kwargs: Any) -> Any: + config = SimpleNamespace(app=app, **kwargs) + configs.append(config) + return config + + class Server: + def __init__(self, config: Any) -> None: + self.config = config + self.should_exit = False + self.ran = threading.Event() + servers.append(self) + + def run(self) -> None: + self.ran.set() + + monkeypatch.setattr(uvicorn, "Config", make_config) + monkeypatch.setattr(uvicorn, "Server", Server) + monkeypatch.setattr(refit_module, "_get_free_port_local", lambda *_args: 12345) + monkeypatch.setattr(refit_module, "_get_node_ip_local", lambda: "10.0.0.1") + config = { + "vllm_cfg": {"async_engine": False, "http_refit_server_port": None}, + "port_range_low": 10000, + "port_range_high": 11000, + } + + with _sparse_refit_receiver(config=config) as receiver: + receiver.setup_api_server = MagicMock() + receiver._setup_vllm_refit_server() + assert configs[0].host == "0.0.0.0" + assert configs[0].port == 12345 + assert servers[0].ran.wait(timeout=1.0) + assert receiver.report_refit_server_base_url() == "http://10.0.0.1:12345" + + relay = MagicMock() + receiver._zmq_refit_server = (relay, "tcp://relay") + receiver._flush_queued_sparse_payloads = MagicMock() + receiver._refit_apply_executor = MagicMock() + receiver.shutdown() + + relay.close.assert_called_once_with() + receiver._flush_queued_sparse_payloads.assert_called_once_with() + receiver._refit_apply_executor.shutdown.assert_called_once_with(wait=True) + assert servers[0].should_exit is True + assert receiver._refit_http_server is None + + +@pytest.mark.asyncio +async def test_async_sparse_refit_post_init_records_worker_locality() -> None: + worker = VllmAsyncGenerationWorkerImpl.__new__(VllmAsyncGenerationWorkerImpl) + worker._sparse_refit_receiver = MagicMock() + worker._mtp_load_from_disk = False + worker.report_device_id_async = AsyncMock(return_value=["0"]) + worker.llm = MagicMock() + worker.llm.collective_rpc = AsyncMock(return_value=["node-0", "node-0"]) + + await worker.post_init_async() + + assert worker.vllm_device_ids == ["0"] + worker._sparse_refit_receiver.set_worker_hostnames.assert_called_once_with( + ["node-0", "node-0"] + ) + assert worker.llm.collective_rpc.await_args_list == [ + call("bind_numa", args=()), + call("report_node_hostname", args=()), + ] + + +def test_async_sparse_refit_exposes_zmq_relay(monkeypatch) -> None: + from nemo_rl.models.generation.vllm import vllm_sparse_refit as refit_module + + server = MagicMock() + server_type = MagicMock(return_value=server) + monkeypatch.setattr(refit_module, "ZmqSparseRefitServer", server_type) + monkeypatch.setattr(refit_module, "_get_free_port_local", lambda *_args: 12345) + monkeypatch.setattr(refit_module, "_get_node_ip_local", lambda: "10.0.0.1") + config = { + "vllm_cfg": { + "async_engine": True, + "zmq_refit_server_port": None, + "http_refit_api_key_env_var": None, + } + } + + with _sparse_refit_receiver(async_engine=True, config=config) as receiver: + assert receiver.start_zmq_sparse_refit_relay(["http://receiver"]) == ( + "tcp://10.0.0.1:12345" + ) + server.start.assert_called_once_with() + + receiver.stop_zmq_sparse_refit_relay() + server.close.assert_called_once_with() + assert receiver._zmq_refit_server is None diff --git a/tests/unit/models/policy/test_megatron_remote_sparse_refit.py b/tests/unit/models/policy/test_megatron_remote_sparse_refit.py new file mode 100644 index 00000000000..c6e33324214 --- /dev/null +++ b/tests/unit/models/policy/test_megatron_remote_sparse_refit.py @@ -0,0 +1,60 @@ +# 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 torch + + +def test_remote_sparse_stream_drains_cuda_before_return(monkeypatch): + from nemo_rl.models.policy.workers.megatron_remote_sparse_refit import ( + MegatronRemoteSparseRefit, + ) + + class Worker: + @staticmethod + def _iter_params_with_optional_kv_scales(): + return iter(()) + + worker = Worker() + remote_refit = object.__new__(MegatronRemoteSparseRefit) + remote_refit._worker = worker + remote_refit._tracker = object() + result = {"payloads": 1, "changed_elements": 2, "total_elements": 3} + events = [] + + def stream(*_args, **_kwargs): + events.append("stream") + return result + + monkeypatch.setattr( + "nemo_rl.models.policy.workers.megatron_remote_sparse_refit." + "stream_sparse_delta_payloads_via_zmq", + stream, + ) + monkeypatch.setattr(torch.cuda, "is_available", lambda: True) + monkeypatch.setattr(torch.cuda, "synchronize", lambda: events.append("sync")) + monkeypatch.setattr(torch.cuda.nvtx, "range_push", lambda _name: None) + monkeypatch.setattr(torch.cuda.nvtx, "range_pop", lambda: None) + + actual = remote_refit.stream( + "zmq", + ["tcp://receiver:5555"], + transfer_id="transfer", + api_key_env_var=None, + timeout_s=1.0, + shard_rank=0, + shard_count=1, + ) + + assert actual is result + assert events == ["stream", "sync"] diff --git a/tests/unit/models/policy/test_megatron_worker.py b/tests/unit/models/policy/test_megatron_worker.py index 7e284ff2f85..793cc52bd90 100644 --- a/tests/unit/models/policy/test_megatron_worker.py +++ b/tests/unit/models/policy/test_megatron_worker.py @@ -18,7 +18,6 @@ from pathlib import Path from types import SimpleNamespace from typing import Optional -from unittest.mock import Mock import numpy as np import pytest @@ -76,67 +75,6 @@ def test_megatron_prepare_for_training_restores_optimizer(): assert restored_devices == ["cuda"] -def test_remote_sparse_stream_drains_cuda_before_return(monkeypatch): - import nemo_rl.models.policy.workers.megatron_policy_worker as worker_module - - worker = object.__new__(worker_module.MegatronPolicyWorkerImpl) - worker.delta_weight_transfer_tracker = object() - worker._iter_params_with_optional_kv_scales = lambda: iter(()) - result = {"payloads": 1, "changed_elements": 2, "total_elements": 3} - events = [] - - def stream(*_args, **_kwargs): - events.append("stream") - return result - - monkeypatch.setattr(worker_module, "stream_sparse_delta_payloads_via_zmq", stream) - monkeypatch.setattr(torch.cuda, "is_available", lambda: True) - monkeypatch.setattr(torch.cuda, "synchronize", lambda: events.append("sync")) - monkeypatch.setattr(torch.cuda.nvtx, "range_push", lambda _name: None) - monkeypatch.setattr(torch.cuda.nvtx, "range_pop", lambda: None) - - actual = worker_module.MegatronPolicyWorkerImpl.stream_remote_sparse_weights( - worker, - "zmq", - ["tcp://receiver:5555"], - transfer_id="transfer", - api_key_env_var=None, - timeout_s=1.0, - shard_rank=0, - shard_count=1, - ) - - assert actual is result - assert events == ["stream", "sync"] - - -def test_refit_param_info_uses_mapping_pp_broadcast(): - from nemo_rl.models.policy.workers.megatron_policy_worker import ( - MegatronPolicyWorkerImpl, - ) - - mapping = SimpleNamespace( - tp_size=2, - ep_size=4, - is_expert=True, - broadcast_obj_from_pp_rank=Mock(return_value=192), - ) - worker = object.__new__(MegatronPolicyWorkerImpl) - worker.dtype = torch.bfloat16 - worker.refit_conversion_tasks = [ - SimpleNamespace( - param_name="decoder.layers.0.mlp.weight", - param_weight=torch.empty(3, 4, dtype=torch.float16), - mapping=mapping, - ) - ] - - assert worker._calculate_refit_param_info() == [ - ("decoder.layers.0.mlp.weight", 192) - ] - mapping.broadcast_obj_from_pp_rank.assert_called_once_with(192) - - def test_set_moe_grad_scale_func_sets_and_clears_on_model_config(): """_set_moe_grad_scale_func should set/clear moe_grad_scale_func on the config.""" from nemo_rl.models.policy.workers.megatron_policy_worker import ( 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 new file mode 100644 index 00000000000..6b5683e90a5 --- /dev/null +++ b/tests/unit/weight_sync/test_vllm_remote_sparse_weight_synchronizer.py @@ -0,0 +1,215 @@ +# 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 unittest.mock import MagicMock, patch + +import pytest + +from nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer import ( + VllmRemoteSparseWeightSynchronizer, +) + + +def _remote_sparse_sync( + mock_ray: MagicMock, + transport: str, + stream_result: list[dict[str, int]] | RuntimeError, +) -> tuple[VllmRemoteSparseWeightSynchronizer, MagicMock, MagicMock]: + init_refs, stream_refs, commit_refs = [MagicMock()], [MagicMock()], [MagicMock()] + policy = MagicMock() + policy.worker_group.workers = [object(), object()] + policy.worker_group.run_all_workers_multiple_data.side_effect = [ + init_refs, + stream_refs, + ] + policy.worker_group.run_all_workers_single_data.return_value = commit_refs + + generation = MagicMock() + generation.worker_group.workers = [object()] + generation.worker_group.run_all_workers_single_data.side_effect = [ + [MagicMock()], + *([[MagicMock()]] if transport == "zmq" else []), + ] + generation.invalidate_kv_cache.return_value = True + + get_results: list[object] = [["http://receiver"]] + if transport == "zmq": + get_results.append(["tcp://relay:19090"]) + get_results.extend([None, stream_result]) + mock_ray.get.side_effect = get_results + + sync = VllmRemoteSparseWeightSynchronizer(policy, generation, transport=transport) + sync.init_communicator() + return sync, policy, generation + + +class TestVllmRemoteSparseWeightSynchronizer: + @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") + def test_init_communicator_requires_receiver_endpoints(self, mock_ray): + policy = MagicMock() + policy.worker_group.workers = [object()] + policy.worker_group.run_all_workers_multiple_data.return_value = [MagicMock()] + generation = MagicMock() + generation.worker_group.workers = [object()] + generation.worker_group.run_all_workers_single_data.return_value = [MagicMock()] + mock_ray.get.return_value = [] + sync = VllmRemoteSparseWeightSynchronizer(policy, generation, transport="s3") + + with pytest.raises(ValueError, match="endpoints are missing"): + sync.init_communicator() + + @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") + def test_shutdown_cancels_pending_work_and_stops_zmq(self, mock_ray): + policy = MagicMock() + generation = MagicMock() + generation.worker_group.workers = [object()] + generation.worker_group.run_all_workers_single_data.return_value = [MagicMock()] + sync = VllmRemoteSparseWeightSynchronizer(policy, generation, transport="zmq") + init_ref, commit_ref = MagicMock(), MagicMock() + sync._baseline_init_refs = [init_ref] + sync._baseline_commit_refs = [commit_ref] + sync._refit_urls = ["http://receiver"] + sync._targets = ["tcp://relay"] + sync._stale = False + + sync.mark_stale() + sync.shutdown() + + assert mock_ray.cancel.call_count == 2 + mock_ray.cancel.assert_any_call(init_ref, force=False) + mock_ray.cancel.assert_any_call(commit_ref, force=False) + generation.worker_group.run_all_workers_single_data.assert_called_once_with( + "stop_zmq_sparse_refit_relay", + run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], + ) + assert sync.is_stale + assert sync._baseline_init_refs is None + assert sync._baseline_commit_refs is None + assert sync._refit_urls == [] + assert sync._targets == [] + + @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") + def test_fails_before_transfer_when_kv_cache_invalidation_fails(self, _mock_ray): + policy = MagicMock() + generation = MagicMock() + generation.invalidate_kv_cache.return_value = False + sync = VllmRemoteSparseWeightSynchronizer(policy, generation, transport="s3") + + with pytest.raises(RuntimeError, match="KV cache invalidation failed"): + sync.sync_weights() + policy.worker_group.run_all_workers_multiple_data.assert_not_called() + + @patch( + "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.flush_vllm_refit_urls" + ) + @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") + def test_initializes_streams_commits_and_updates_baseline( + self, mock_ray, flush, capsys + ): + sync, policy, generation = _remote_sparse_sync( + mock_ray, + "zmq", + [{"payloads": 3, "changed_elements": 3, "total_elements": 100}], + ) + flush.return_value = [ + { + "verification_candidates": 4, + "verification_samples": 4, + "verification_exact_mismatches": 1, + "verification_mismatches": 0, + "verification_abs_sum": 1e-9, + "verification_max_abs": 1e-9, + } + ] + metrics = sync.sync_weights() + + assert [ + entry.args[0] + for entry in policy.worker_group.run_all_workers_multiple_data.call_args_list + ] == ["init_remote_sparse_delta_baseline", "stream_remote_sparse_weights"] + assert [ + entry.args[0] + for entry in generation.worker_group.run_all_workers_single_data.call_args_list + ] == ["report_refit_server_base_url", "start_zmq_sparse_refit_relay"] + generation.worker_group.run_all_workers_single_data.assert_any_call( + "start_zmq_sparse_refit_relay", + run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], + refit_urls=["http://receiver"], + ) + flush.assert_called_once_with( + ["http://receiver"], api_key_env_var=None, timeout_s=600.0 + ) + policy.worker_group.run_all_workers_single_data.assert_called_once_with( + "finish_remote_sparse_delta_sync", succeeded=True + ) + assert ( + "REFIT_ZMQ_DELTA_CHANGE changed_elements=3 total_elements=100 " + "changed_pct=3" in capsys.readouterr().out + ) + assert metrics["delta/changed_pct"] == 3.0 + assert metrics["delta_verify/candidates"] == 4.0 + assert metrics["delta_verify/samples"] == 4.0 + assert metrics["delta_verify/exact_mismatches"] == 1.0 + assert metrics["delta_verify/mismatches"] == 0.0 + assert metrics["delta_verify/mean_abs"] == 2.5e-10 + assert metrics["delta_verify/max_abs"] == 1e-9 + assert metrics["transfer/payloads"] == 3.0 + assert not sync.is_stale + + @patch( + "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.flush_vllm_refit_urls" + ) + @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") + def test_sample_mismatch_does_not_commit_baseline(self, mock_ray, flush): + sync, policy, _ = _remote_sparse_sync( + mock_ray, + "zmq", + [{"payloads": 3, "changed_elements": 3, "total_elements": 100}], + ) + flush.return_value = [ + { + "verification_samples": 4, + "verification_mismatches": 1, + "verification_abs_sum": 0.5, + "verification_max_abs": 0.5, + } + ] + + with pytest.raises(RuntimeError, match="1 mismatched deltas out of 4"): + sync.sync_weights() + + policy.worker_group.run_all_workers_single_data.assert_called_once_with( + "finish_remote_sparse_delta_sync", succeeded=False + ) + + @patch( + "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.flush_vllm_refit_urls" + ) + @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") + def test_failure_drains_receivers_without_committing_baseline( + self, mock_ray, flush + ): + sync, policy, _ = _remote_sparse_sync( + mock_ray, "s3", RuntimeError("stream failed") + ) + + with pytest.raises(RuntimeError, match="stream failed"): + sync.sync_weights() + + flush.assert_called_once_with( + ["http://receiver"], api_key_env_var=None, timeout_s=60.0 + ) + policy.worker_group.run_all_workers_single_data.assert_called_once_with( + "finish_remote_sparse_delta_sync", succeeded=False + ) diff --git a/tests/unit/weight_sync/test_weight_synchronizer.py b/tests/unit/weight_sync/test_weight_synchronizer.py index 99b593ec9d9..edec20b6c0a 100644 --- a/tests/unit/weight_sync/test_weight_synchronizer.py +++ b/tests/unit/weight_sync/test_weight_synchronizer.py @@ -34,9 +34,6 @@ from nemo_rl.weight_sync.ipc_weight_synchronizer import ( IPCWeightSynchronizer, ) -from nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer import ( - VllmRemoteSparseWeightSynchronizer, -) # --------------------------------------------------------------------------- # Helpers @@ -79,25 +76,6 @@ def _mock_cluster(world_size=4, ip="127.0.0.1", port=29500): return cluster -def _remote_sparse_sync( - mock_ray: MagicMock, - transport: str, - stream_result: list[dict[str, int]] | RuntimeError, -) -> tuple[VllmRemoteSparseWeightSynchronizer, MagicMock, MagicMock]: - policy = MagicMock() - policy.init_remote_sparse_delta_baseline.return_value = [MagicMock()] - policy.stream_remote_sparse_weights.return_value = [MagicMock()] - policy.finish_remote_sparse_delta_sync.return_value = [MagicMock()] - generation = MagicMock() - generation.report_refit_server_base_urls.return_value = ["http://receiver"] - generation.start_zmq_sparse_refit_relays.return_value = ["tcp://relay:19090"] - generation.invalidate_kv_cache.return_value = True - mock_ray.get.side_effect = [None, stream_result] - sync = VllmRemoteSparseWeightSynchronizer(policy, generation, transport=transport) - sync.init_communicator() - return sync, policy, generation - - # --------------------------------------------------------------------------- # WeightSynchronizer ABC contract # --------------------------------------------------------------------------- @@ -238,143 +216,6 @@ def test_zero_env_ratio_raises(self, mock_ray, monkeypatch): sync._compute_buffer_size() -class TestVllmRemoteSparseWeightSynchronizer: - def test_init_communicator_requires_receiver_endpoints(self): - policy = MagicMock() - generation = MagicMock() - generation.report_refit_server_base_urls.return_value = [] - sync = VllmRemoteSparseWeightSynchronizer(policy, generation, transport="s3") - - with pytest.raises(ValueError, match="endpoints are missing"): - sync.init_communicator() - - @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") - def test_shutdown_cancels_pending_work_and_stops_zmq(self, mock_ray): - policy = MagicMock() - generation = MagicMock() - sync = VllmRemoteSparseWeightSynchronizer(policy, generation, transport="zmq") - init_ref, commit_ref = MagicMock(), MagicMock() - sync._baseline_init_refs = [init_ref] - sync._baseline_commit_refs = [commit_ref] - sync._refit_urls = ["http://receiver"] - sync._targets = ["tcp://relay"] - sync._stale = False - - sync.mark_stale() - sync.shutdown() - - assert mock_ray.cancel.call_count == 2 - mock_ray.cancel.assert_any_call(init_ref, force=False) - mock_ray.cancel.assert_any_call(commit_ref, force=False) - generation.stop_zmq_sparse_refit_relays.assert_called_once_with() - assert sync.is_stale - assert sync._baseline_init_refs is None - assert sync._baseline_commit_refs is None - assert sync._refit_urls == [] - assert sync._targets == [] - - @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") - def test_fails_before_transfer_when_kv_cache_invalidation_fails(self, _mock_ray): - policy = MagicMock() - generation = MagicMock() - generation.invalidate_kv_cache.return_value = False - sync = VllmRemoteSparseWeightSynchronizer(policy, generation, transport="s3") - - with pytest.raises(RuntimeError, match="KV cache invalidation failed"): - sync.sync_weights() - policy.stream_remote_sparse_weights.assert_not_called() - - @patch( - "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.flush_vllm_refit_urls" - ) - @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") - def test_initializes_streams_commits_and_updates_baseline( - self, mock_ray, flush, capsys - ): - sync, policy, generation = _remote_sparse_sync( - mock_ray, - "zmq", - [{"payloads": 3, "changed_elements": 3, "total_elements": 100}], - ) - flush.return_value = [ - { - "verification_candidates": 4, - "verification_samples": 4, - "verification_exact_mismatches": 1, - "verification_mismatches": 0, - "verification_abs_sum": 1e-9, - "verification_max_abs": 1e-9, - } - ] - metrics = sync.sync_weights() - - policy.init_remote_sparse_delta_baseline.assert_called_once_with("zmq") - generation.start_zmq_sparse_refit_relays.assert_called_once_with( - ["http://receiver"] - ) - policy.stream_remote_sparse_weights.assert_called_once() - flush.assert_called_once_with( - ["http://receiver"], api_key_env_var=None, timeout_s=600.0 - ) - policy.finish_remote_sparse_delta_sync.assert_called_once_with(True) - assert ( - "REFIT_ZMQ_DELTA_CHANGE changed_elements=3 total_elements=100 " - "changed_pct=3" in capsys.readouterr().out - ) - assert metrics["delta/changed_pct"] == 3.0 - assert metrics["delta_verify/candidates"] == 4.0 - assert metrics["delta_verify/samples"] == 4.0 - assert metrics["delta_verify/exact_mismatches"] == 1.0 - assert metrics["delta_verify/mismatches"] == 0.0 - assert metrics["delta_verify/mean_abs"] == 2.5e-10 - assert metrics["delta_verify/max_abs"] == 1e-9 - assert metrics["transfer/payloads"] == 3.0 - assert not sync.is_stale - - @patch( - "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.flush_vllm_refit_urls" - ) - @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") - def test_sample_mismatch_does_not_commit_baseline(self, mock_ray, flush): - sync, policy, _ = _remote_sparse_sync( - mock_ray, - "zmq", - [{"payloads": 3, "changed_elements": 3, "total_elements": 100}], - ) - flush.return_value = [ - { - "verification_samples": 4, - "verification_mismatches": 1, - "verification_abs_sum": 0.5, - "verification_max_abs": 0.5, - } - ] - - with pytest.raises(RuntimeError, match="1 mismatched deltas out of 4"): - sync.sync_weights() - - policy.finish_remote_sparse_delta_sync.assert_called_once_with(False) - - @patch( - "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.flush_vllm_refit_urls" - ) - @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") - def test_failure_drains_receivers_without_committing_baseline( - self, mock_ray, flush - ): - sync, policy, _ = _remote_sparse_sync( - mock_ray, "s3", RuntimeError("stream failed") - ) - - with pytest.raises(RuntimeError, match="stream failed"): - sync.sync_weights() - - flush.assert_called_once_with( - ["http://receiver"], api_key_env_var=None, timeout_s=60.0 - ) - policy.finish_remote_sparse_delta_sync.assert_called_once_with(False) - - # --------------------------------------------------------------------------- # HTTPWeightSynchronizer # --------------------------------------------------------------------------- From 983f82fac85e695c26ec5ebb66f6e551936de401 Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Fri, 10 Jul 2026 14:10:06 -0700 Subject: [PATCH 05/18] Add calculator Signed-off-by: Hollow Man --- docs/design-docs/sparse-delta-refit.md | 410 ++++++++++++++++++++----- pyrefly.toml | 1 + tools/refit_bandwidth_calculator.py | 319 +++++++++++++++++++ 3 files changed, 660 insertions(+), 70 deletions(-) create mode 100644 tools/refit_bandwidth_calculator.py diff --git a/docs/design-docs/sparse-delta-refit.md b/docs/design-docs/sparse-delta-refit.md index 743a9fc31ea..96d01345388 100644 --- a/docs/design-docs/sparse-delta-refit.md +++ b/docs/design-docs/sparse-delta-refit.md @@ -1,78 +1,348 @@ # Remote Sparse-Delta vLLM Refit -For non-colocated Megatron policy workers and sync vLLM workers that share the -same checkpoint. Policy workers keep a CPU baseline and stream zstd-compressed -sparse deltas through either S3 or ZeroMQ. Both transports share export, -encoding, compression, backpressure, receiver apply, and transactional baseline -commit logic. -Payload checksums and transfer-scoped IDs make HTTP retries idempotent. Policy -workers commit their baselines only after every receiver flush succeeds. -The implementation is opt-in. Core policy and vLLM workers contain only lazy -delegates; export/encoding, transport, receiver queuing, and sparse placement -live in dedicated modules and are not initialized by existing IPC, HTTP, or -NCCL refit paths. +Remote sparse-delta refit updates non-colocated vLLM workers without sending a +full checkpoint after every optimizer step. Megatron workers export Hugging +Face (HF) weights, compare them with a sharded CPU baseline, and send changed +locations and deltas through S3 or ZeroMQ. vLLM maps those HF coordinates into +its local TP and EP layouts and applies the updates in place. -On a fresh run, generation starts from the shared checkpoint while policy -workers build the CPU baseline asynchronously; the first transfer follows the -first optimizer step. Resumed runs synchronize before generation. +The feature is opt-in. Its synchronizer, codec, transports, receiver queue, and +placement engine are separate from existing NCCL, CUDA IPC, and packed refit +paths. -## Config +## Supported scope + +Remote sparse refit requires: + +- a non-colocated Megatron policy and vLLM generation backend; +- the same initial HF checkpoint on both clusters; +- BF16 or FP16, unquantized rollout weights; +- `kv_cache_dtype: auto`; and +- a `delta_compression` configuration. + +Configuration validation rejects FP8 weights, FP8 KV-cache scales, +`quant_cfg`, `real_quant`, colocated inference, and non-Megatron policies. +Synchronous and asynchronous vLLM engines are supported, but the weight-version +transition is synchronous: generation pauses until all payloads are applied and +the global flush completes. + +## Architecture + +```mermaid +flowchart LR + subgraph P["Megatron policy cluster"] + B["Megatron Bridge HF export"] + C["Sharded CPU or mmap baseline"] + E["Compare, encode, and compress"] + B --> C + B --> E + C --> E + end + + S["S3 object and HTTP manifest"] + Z["ZeroMQ relay"] + + subgraph G["vLLM generation cluster"] + H["HTTP receiver"] + Q["Bounded FIFO apply queue"] + A["Sparse placement and apply"] + H --> Q --> A + end + + E -->|S3| S --> H + E -->|ZeroMQ| Z --> H +``` + +*Figure 1. S3 and ZeroMQ share the exporter, codec, receiver, placement engine, +and commit protocol.* + +| Responsibility | Implementation | +|---|---| +| Coordinate one transfer and commit | [`vllm_remote_sparse_weight_synchronizer.py`](../../nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py) | +| Adapt Megatron workers | [`megatron_remote_sparse_refit.py`](../../nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py) | +| Track baselines and encode deltas | [`weight_transfer_sparse_codec.py`](../../nemo_rl/utils/weight_transfer_sparse_codec.py) | +| Run the shared pipeline and S3 transport | [`weight_transfer_remote_sparse.py`](../../nemo_rl/utils/weight_transfer_remote_sparse.py) | +| Run the ZeroMQ transport and relay | [`weight_transfer_zmq.py`](../../nemo_rl/utils/weight_transfer_zmq.py) | +| Queue receiver work and expose endpoints | [`vllm_sparse_refit.py`](../../nemo_rl/models/generation/vllm/vllm_sparse_refit.py) | +| Map HF coordinates into vLLM tensors | [`vllm_sparse_delta.py`](../../nemo_rl/models/generation/vllm/vllm_sparse_delta.py) | + +## Refit protocol + +### Initialize the baseline + +`VllmRemoteSparseWeightSynchronizer.init_communicator()` starts baseline +construction on every policy worker and then discovers the vLLM HTTP endpoints. +Each worker participates in the HF export but stores only chunks assigned by +`chunk_index % shard_count`. The baseline is therefore sharded across policy +workers. It uses file-backed `torch.from_file` tensors by default; +`NRL_REFIT_BASELINE_IN_MEMORY=1` keeps it in RAM. + +On a fresh run, vLLM already holds the shared checkpoint. Baseline construction +starts early and can overlap initial generation, so the redundant initial full +sync is skipped. A resumed run performs a refit before generation because the +training checkpoint may be newer than the rollout checkpoint. + +### Export and encode deltas + +Every refit still traverses `MegatronBridge.export_hf_weights()`. Sharding the +baseline avoids duplicate baseline storage and payload production, but it does +not remove Bridge export or the full CPU comparison. Low changed density mainly +reduces encoded and transferred bytes. + +For each assigned chunk, `DeltaCompressionTracker` finds changed flat +locations and encodes deltas in the configured dtype. The pending baseline is +updated to the value the receiver will hold after wire-dtype rounding: + +```text +expected = previous_baseline + delta.to(baseline_dtype) +``` + +The producer overlaps export, encoding, `torch.save` serialization, zstd level +1 compression, and transfer with bounded executors. Source baselines do not +commit until the entire transfer succeeds. + +### Transfer and apply + +| Behavior | S3 | ZeroMQ | +|---|---|---| +| Value plane | AWS CRT `PUT_OBJECT` | DEALER to ROUTER relay | +| Receiver notification | HTTP object manifest | Relay HTTP fanout | +| Retry identity | Object key and checksum | Transfer, producer, payload IDs, and checksum | +| Lifetime | Delete after all receivers respond | No persistent object | + +S3 uses 64 MiB multipart parts, a 2 GiB client memory limit, and a 10 Gbps CRT +throughput target. ZeroMQ assigns each producer to one relay; that relay fans +the compressed payload out to every generation replica. Both transports use +the same receiver endpoints and checksum validation. + +The receiver deduplicates payload identities, batches them in a bounded FIFO +queue, and applies batches on one worker thread. When all vLLM ranks share a +node, payloads are staged under `/dev/shm` and passed to collective RPC by file +path. Otherwise, each serialized payload is sent through collective RPC. + +The final `/nemo-rl/refit/flush` drains the queue, synchronizes CUDA, and checks +optional delta samples. Only then does the source commit pending baseline +updates in background CPU threads. + +> **Failure boundary:** source baseline commit is transactional, but receiver +> updates are in place and are not rolled back. If a transfer fails after a +> receiver accepts any payload, reload that receiver from a known-good weight +> version before retrying. + +## Payload and placement + +Each serialized payload is: + +```text +(packed_location_bytes, packed_delta_values, tensor_metadata) +``` + +Contiguous locations use a range encoding. Other sorted locations are +delta-encoded into the smallest lossless unsigned width among 16, 32, and 64 +bits. Metadata carries the HF name and shape, value offsets, location encoding, +and optional verification samples. + +HF coordinates are the canonical wire format because Megatron Bridge already +defines the training-to-HF mapping while vLLM owns a different packed and +sharded layout. `_SparseDeltaTargetPlan` converts source coordinates to local +vLLM indices without materializing a dense HF tensor. Plans cover identity and +single-dimension shards, packed QKV, merged gate/up projections, fused MoE +experts, and Mamba2 layouts. + +`_SparseDeltaTargetPlan(target=None)` means a valid tensor is absent from the +current rank. A `None` plan means the layout is unsupported; the receiver fails +the payload before applying it. There is no dense fallback for an unknown +layout. + +## Configuration + +Configure the feature under `policy.generation`: ```yaml -backend: vllm -colocated: {enabled: false} -refit_transport: vllm_s3_sparse -delta_compression: - dtype: bf16 - sparse_bucket_size_bytes: 268435456 -vllm_cfg: - async_engine: false - http_refit_server_port: 8081 - http_refit_api_key_env_var: NRL_REFIT_API_KEY +policy: + generation: + backend: vllm + refit_transport: vllm_s3_sparse # or vllm_zmq_sparse + delta_compression: + dtype: bf16 + sparse_bucket_size_bytes: 268435456 + colocated: + enabled: false + vllm_cfg: + async_engine: false # true is also supported + precision: bfloat16 + kv_cache_dtype: auto + http_refit_api_key_env_var: NRL_REFIT_API_KEY + http_refit_server_port: 8081 + zmq_refit_server_port: null +``` + +S3 requires `NRL_REFIT_S3_BUCKET`; region and key prefix default to +`us-east-1` and `nemo-rl-refit`. ZeroMQ requires routable TCP access to the +relay port. The HTTP and ZeroMQ servers are plaintext, so use a trusted or +encrypted network. When `http_refit_api_key_env_var` is set, the named variable +must contain the same nonempty token on producers and receivers. + +| Control | Default | +|---|---:| +| `NRL_REFIT_S3_EXPORT_CHUNK_BYTES` | 256 MiB | +| `NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES` | 1 GiB | +| `NRL_REFIT_{S3,ZMQ}_ENCODE_WORKERS` | 2-8 from CPU count | +| `NRL_REFIT_S3_UPLOAD_WORKERS` | 4-32 from CPU count | +| `NRL_REFIT_ZMQ_SEND_WORKERS` | 4 | +| `NRL_REFIT_ZMQ_RELAY_PAYLOAD_WORKERS` | 16 | +| `NRL_REFIT_ZMQ_RELAY_FANOUT_WORKERS` | 8-32 from replica count | +| `NRL_REFIT_APPLY_QUEUE_DEPTH` / `NRL_REFIT_APPLY_BATCH_SIZE` | 2 / 8 | +| `NRL_REFIT_{S3,ZMQ}_ZSTD_THREADS` | 0 | +| `NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD` | 0 | + +Export chunks are also capped by `sparse_bucket_size_bytes` and the packed +tensor limit. Increase one concurrency control at a time; excessive parallelism +can move the bottleneck into host memory, collective export, relay fanout, or +receiver apply. + +## Metrics and profiling + +| Signal | Meaning | +|---|---| +| `REFIT_BASELINE_INIT` | Baseline export and snapshot time | +| `REFIT_{S3,ZMQ}_TIMING` | Producer wall time, stage service time, payloads, bytes, and changed density | +| `REFIT_{S3,ZMQ}_DELTA_CHANGE` | Global changed and total element counts | +| `REFIT_RECEIVER_TIMING` | Receiver batches, apply time, and verification counts | +| `REFIT_{S3,ZMQ}_DELTA_VERIFY` | Sampled transmitted-delta accuracy | +| `REFIT_{S3,ZMQ}_GLOBAL_COMMIT` | Successful transfer flush | + +`total_s` is producer wall time. Stage fields such as `encode_s`, `s3_put_s`, +and `zmq_send_s` are sums across concurrent tasks and can exceed `total_s`; do +not add them as serial phases. + +The synchronizer returns metrics under `refit/delta/*`, +`refit/delta_verify/*`, and `refit/transfer/*` when GRPO logs them. These are +available to W&B and other configured loggers. End-to-end refit latency is +reported as `timing/train/prepare_for_generation/transfer_and_update_weights`. + +For Nsight Systems, use the existing baseline, policy stream, and vLLM +sparse-apply NVTX ranges. Producer and receiver thread names begin with +`nrl-refit-`, `nrl-zmq-`, or `nrl-vllm-sparse-refit`. + +## Development and validation + +Keep transport changes behind the shared `stream_sparse_delta_payloads()` +pipeline. A transport should provide payload delivery and timing only; it must +not duplicate the baseline tracker, codec, receiver queue, or placement logic. +Retries must preserve payload identity and bytes, fan out to every required +replica, and require a successful global flush before baseline commit. + +For a new vLLM layout, derive shard ownership and offsets from vLLM module +attributes or its loader contract. Validate every source shape and target +capacity. Unit tests must exercise nonzero TP ranks, replicated KV heads, local +and remote experts, uneven shapes, contiguous ranges, and explicit locations. +In-range but incorrect `index_add_` locations silently corrupt weights, so test +the exact mapped indices and values. + +Codec changes must update encoder and decoder together, preserve 64-bit-safe +locations, and retain wire-dtype rounding in pending baseline updates. Receiver +changes must preserve FIFO application, bounded memory, deferred-error +propagation, flush, CUDA synchronization, and clean shutdown. + +Run the focused suite: + +```bash +uv run pytest -q \ + tests/unit/utils/test_weight_transfer_remote_sparse.py \ + tests/unit/models/policy/test_megatron_remote_sparse_refit.py \ + tests/unit/models/generation/test_vllm_sparse_refit.py \ + tests/unit/models/generation/test_vllm_sparse_delta.py \ + tests/unit/weight_sync/test_vllm_remote_sparse_weight_synchronizer.py \ + tests/unit/tools/test_refit_bandwidth_calculator.py + +uv run ruff check \ + nemo_rl/utils/weight_transfer_{remote_sparse,sparse_codec,zmq}.py \ + nemo_rl/models/generation/vllm/vllm_{sparse_refit,sparse_delta}.py \ + nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py \ + tools/refit_bandwidth_calculator.py ``` -Use `refit_transport: vllm_zmq_sparse` and set -`vllm_cfg.zmq_refit_server_port` when Kubernetes needs a stable ZeroMQ target -port. The ZeroMQ service must route TCP traffic from policy workers to the vLLM -relay workers; each relay fans a payload out to every HTTP refit endpoint. On a -flat cluster network, the dynamically reported worker IP can be used directly; -a service mesh is not required. - -Remote sparse refit requires `kv_cache_dtype: auto`; FP8 KV-cache scale sync is not -supported. Receiver tensors must have a direct QKV, MoE, Mamba, or generic TP -placement plan; transformed and FP8 weights fail before any delta is applied. -The HTTP endpoints and ZeroMQ producer relay use -`http_refit_api_key_env_var` auth when configured. - -For S3, set `NRL_REFIT_S3_BUCKET` and, when needed, -`NRL_REFIT_S3_REGION` or `NRL_REFIT_S3_PREFIX`. AWS CRT performs multipart -transfer automatically. Tune export, encode, and transfer concurrency with -`NRL_REFIT_S3_EXPORT_CHUNK_BYTES`, `NRL_REFIT_S3_ENCODE_WORKERS`, and -`NRL_REFIT_S3_UPLOAD_WORKERS`. - -For ZeroMQ, tune the same stages with `NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES`, -`NRL_REFIT_ZMQ_ENCODE_WORKERS`, and `NRL_REFIT_ZMQ_SEND_WORKERS`. Relay -concurrency is controlled by `NRL_REFIT_ZMQ_RELAY_PAYLOAD_WORKERS` and -`NRL_REFIT_ZMQ_RELAY_FANOUT_WORKERS`; fanout defaults to 32 workers to preserve -HTTP keep-alive reuse and avoid receiver contention. Track `REFIT_S3_TIMING` or -`REFIT_ZMQ_TIMING`, `REFIT_RECEIVER_TIMING`, and `REFIT_*_GLOBAL_COMMIT` in -cluster runs. - -Set `NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD` to a small positive value to verify -deterministic transmitted-delta samples after placement. Each receiver snapshots -only those target elements before apply and compares `post - pre` with the -placement- and dtype-adjusted transmitted delta, so an existing absolute weight -offset cannot contaminate the metric. `REFIT_*_DELTA_VERIFY` reports candidate -and applied sample counts, exact mismatches, tolerance-gated mismatches -(`rtol=1e-6`, `atol=1e-8`), and mean/max absolute delta error. A gated mismatch -aborts the transaction before baseline commit. Successful commits advance the -CPU baseline to the quantized value applied by the receiver, so later deltas -compensate compression residuals instead of accumulating drift. - -`REFIT_*_DELTA_CHANGE` reports changed and total exported element counts plus -their model-wide percentage. The codec accumulates these counters while it is -already finding sparse locations; it does not perform another tensor scan. -When a training logger is enabled, the same values are emitted to W&B and -TensorBoard under `refit/delta/*`, `refit/delta_verify/*`, and -`refit/transfer/*`. End-to-end refit latency remains available as -`timing/train/prepare_for_generation/transfer_and_update_weights`. +On the target topology, verify the exact commit, image digest, and checkpoint +revision; run fresh and resumed starts; compare at least two balanced +repetitions with an equivalent NCCL or full control; and require the requested +changed density, one global commit, no traceback, and zero sampled mismatches. +After failure injection, confirm the source baseline does not commit and reload +the receiver before retrying. + +## Refit bandwidth calculator + +[`refit_bandwidth_calculator.py`](../../tools/refit_bandwidth_calculator.py) is a +benchmark-specific estimator for the current S3 and ZeroMQ implementation. It +is not a general fabric or topology model. + +The sparse side embeds July 2026 end-to-end fits from 32 GB300 sender GPUs in +`us-east-2` to 64 H100 receiver GPUs in `us-east-1`. The measured checkpoints +span 63.2-1121.0 GB of indexed BF16 weights. With `S` as decimal TB, the +embedded latency fits are: + +| Transport | Payload | 3% changed | 5% changed | +|---|---|---:|---:| +| S3 | raw | `2.370 + 116.138S` | `2.646 + 183.431S` | +| S3 | zstd | `84.746S` | `134.041S` | +| ZeroMQ | raw | `264.753S` | `428.008S` | +| ZeroMQ | zstd | `6.517 + 73.621S` | `8.165 + 157.339S` | + +For an arbitrary positive changed density `d`, the calculator evaluates `T3` +and `T5` from this table, computes `p = log(T5 / T3) / log(5 / 3)`, and returns +`T3 * (d / 3)^p`. Inputs outside 3-5% are accepted but are extrapolations. The +wire estimate is `model_size_gb * d / 100 * multiplier`, where the raw and zstd +multipliers are 2.0 and 0.74. + +The NCCL side uses these measured generation-EP H100 refit envelopes on 400 +Gbps/rank InfiniBand: + +| Indexed BF16 | NCCL refit envelope | +|---:|---:| +| 63.2 GB | 0.84-1.60 s | +| 247.2 GB | 1.46-1.74 s | +| 470.2 GB | 2.31-2.73 s | +| 1342.0 GB | 3.27-3.46 s | + +The calculator interpolates these anchors in log model-size space, then +projects the full measured NCCL latency onto the candidate Ethernet rate: + +```text +T_ethernet = T_H100_IB * 400 / candidate_ethernet_gbps +``` + +`--candidate-ethernet-gbps` is raw bandwidth per rank. It changes only the NCCL +projection; it does not rescale the measured S3 or ZeroMQ fit. Do not pass +aggregate node or cluster bandwidth. + +```bash +uv run python tools/refit_bandwidth_calculator.py \ + --model-size-gb 247.2 \ + --changed-pct 3 \ + --compression zstd \ + --candidate-ethernet-gbps 25 +``` + +The output reports the original H100 IB envelope, projected NCCL latency, +estimated sparse latency and wire bytes, and an Ethernet crossover range. Below +the lower crossover, sparse refit beats the complete NCCL envelope; above the +upper crossover, NCCL wins; between them, the measured NCCL range does not give +one winner. `--json` emits the same data for scripts. + +The production transport currently applies zstd level 1 to every payload. +`--compression raw` selects a historical uncompressed benchmark fit for +analysis; it is not a runtime switch for the current transport. Treat model +sizes outside the measured range, changed densities outside 3-5%, and different +parallel mappings as experiment targets rather than performance claims. + +## Failure guide + +| Symptom | Action | +|---|---| +| Baseline is missing a tensor | Check baseline completion, checkpoint equality, and Bridge name mappings. | +| No refit endpoint is found | Check worker startup, fixed ports, routing, and network policy. | +| No direct target plan exists | Add and unit-test the layout; do not silently fall back to dense loading. | +| A payload ID is reused with different bytes | Start a new transfer or resend the original payload unchanged. | +| Changed percentage rises unexpectedly | Correlate `DELTA_CHANGE` with `GLOBAL_COMMIT` and baseline commit completion. | +| Apply queue stalls | Inspect receiver timing and reduce source or relay concurrency. | +| A transfer fails after payload acceptance | Reload the receiver from a known-good checkpoint before retrying. | diff --git a/pyrefly.toml b/pyrefly.toml index 4ee30d2634b..7e89fff230f 100644 --- a/pyrefly.toml +++ b/pyrefly.toml @@ -207,6 +207,7 @@ project-includes = [ "nemo_rl/utils/weight_transfer_remote_sparse.py", "nemo_rl/utils/weight_transfer_sparse_codec.py", "nemo_rl/utils/weight_transfer_zmq.py", + "tools/refit_bandwidth_calculator.py", "tools/model_diagnostics/1.max_model_len_respected.py", "tools/model_diagnostics/2.long_generation_decode_vs_prefill.py", "tools/model_diagnostics/3.check_and_reinit_hf_model_embeddings_untrained.py", diff --git a/tools/refit_bandwidth_calculator.py b/tools/refit_bandwidth_calculator.py new file mode 100644 index 00000000000..be82173a0b0 --- /dev/null +++ b/tools/refit_bandwidth_calculator.py @@ -0,0 +1,319 @@ +# 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. + +"""Estimate when measured sparse refit beats NCCL over Ethernet. + +This is a benchmark-specific estimator. Sparse latency is fitted from the July +2026 S3 and ZeroMQ benchmarks, and NCCL latency is interpolated from measured +H100 reshard results on 400 Gbps/rank InfiniBand using hierarchical API. +``--candidate-ethernet-gbps`` projects that NCCL reference onto a raw +per-rank Ethernet rate. +""" + +import argparse +import json +import math +from dataclasses import asdict, dataclass +from itertools import pairwise +from typing import Literal + +Transport = Literal["s3", "zmq"] +Compression = Literal["raw", "zstd"] + +_REFERENCE_IB_GBPS = 400.0 +_CALIBRATED_DENSITIES = (3.0, 5.0) +_CALIBRATED_MODEL_SIZE_GB = (63.2, 1121.0) + + +@dataclass(frozen=True) +class _SparseFit: + intercept_s: float + seconds_per_tb: float + + +# T(S) = intercept_s + seconds_per_tb * S for decimal TB of indexed BF16. +_SPARSE_FITS: dict[tuple[Transport, Compression, float], _SparseFit] = { + ("s3", "raw", 3.0): _SparseFit(2.370, 116.138), + ("s3", "zstd", 3.0): _SparseFit(0.000, 84.746), + ("s3", "raw", 5.0): _SparseFit(2.646, 183.431), + ("s3", "zstd", 5.0): _SparseFit(0.000, 134.041), + ("zmq", "raw", 3.0): _SparseFit(0.000, 264.753), + ("zmq", "zstd", 3.0): _SparseFit(6.517, 73.621), + ("zmq", "raw", 5.0): _SparseFit(0.000, 428.008), + ("zmq", "zstd", 5.0): _SparseFit(8.165, 157.339), +} + +# (indexed model GB, low seconds, high seconds), with generation EP enabled. +_NCCL_ANCHORS = ( + (63.2, 0.84, 1.60), + (247.2, 1.46, 1.74), + (470.2, 2.31, 2.73), + (1342.0, 3.27, 3.46), +) + +# Approximate serialized bytes / changed BF16 bytes. +_WIRE_MULTIPLIER: dict[Compression, float] = {"raw": 2.0, "zstd": 0.74} + + +@dataclass(frozen=True) +class Estimate: + """One sparse transport compared with the NCCL reference.""" + + transport: Transport + compression: Compression + model_size_gb: float + changed_pct: float + sparse_seconds: float + approximate_wire_gb: float + nccl_ib_low_s: float + nccl_ib_high_s: float + candidate_ethernet_gbps: float | None + nccl_ethernet_low_s: float | None + nccl_ethernet_high_s: float | None + break_even_ethernet_low_gbps: float + break_even_ethernet_high_gbps: float + candidate_winner: str | None + model_size_extrapolated: bool + density_extrapolated: bool + + +def predict_sparse_seconds( + model_size_gb: float, + changed_pct: float, + *, + transport: Transport, + compression: Compression, +) -> float: + """Interpolate or extrapolate sparse latency from the 3% and 5% fits.""" + if model_size_gb <= 0 or changed_pct <= 0: + raise ValueError("model_size_gb and changed_pct must be positive") + + size_tb = model_size_gb / 1000.0 + density_low, density_high = _CALIBRATED_DENSITIES + latency_low = _fit_latency(size_tb, transport, compression, density_low) + latency_high = _fit_latency(size_tb, transport, compression, density_high) + exponent = math.log(latency_high / latency_low) / math.log( + density_high / density_low + ) + return latency_low * (changed_pct / density_low) ** exponent + + +def _fit_latency( + size_tb: float, + transport: Transport, + compression: Compression, + changed_pct: float, +) -> float: + fit = _SPARSE_FITS[(transport, compression, changed_pct)] + return fit.intercept_s + fit.seconds_per_tb * size_tb + + +def _nccl_reference(model_size_gb: float) -> tuple[float, float, bool]: + if model_size_gb <= 0: + raise ValueError("model_size_gb must be positive") + + extrapolated = not _NCCL_ANCHORS[0][0] <= model_size_gb <= _NCCL_ANCHORS[-1][0] + if model_size_gb <= _NCCL_ANCHORS[0][0]: + left, right = _NCCL_ANCHORS[:2] + elif model_size_gb >= _NCCL_ANCHORS[-1][0]: + left, right = _NCCL_ANCHORS[-2:] + else: + left, right = _NCCL_ANCHORS[:2] + for candidate_left, candidate_right in pairwise(_NCCL_ANCHORS): + if candidate_left[0] <= model_size_gb <= candidate_right[0]: + left, right = candidate_left, candidate_right + break + + position = math.log(model_size_gb / left[0]) / math.log(right[0] / left[0]) + low_s = left[1] + position * (right[1] - left[1]) + high_s = left[2] + position * (right[2] - left[2]) + return max(0.001, low_s), max(0.001, high_s), extrapolated + + +def estimate( + *, + model_size_gb: float, + changed_pct: float, + transport: Transport, + compression: Compression = "zstd", + candidate_ethernet_gbps: float | None = None, +) -> Estimate: + """Estimate sparse latency and the NCCL-over-Ethernet crossover.""" + if candidate_ethernet_gbps is not None and candidate_ethernet_gbps <= 0: + raise ValueError("candidate_ethernet_gbps must be positive") + + sparse_seconds = predict_sparse_seconds( + model_size_gb, + changed_pct, + transport=transport, + compression=compression, + ) + nccl_low_s, nccl_high_s, model_size_extrapolated = _nccl_reference(model_size_gb) + break_even_low = _REFERENCE_IB_GBPS * nccl_low_s / sparse_seconds + break_even_high = _REFERENCE_IB_GBPS * nccl_high_s / sparse_seconds + + candidate_low = candidate_high = None + winner = None + if candidate_ethernet_gbps is not None: + scale = _REFERENCE_IB_GBPS / candidate_ethernet_gbps + candidate_low = nccl_low_s * scale + candidate_high = nccl_high_s * scale + if sparse_seconds < candidate_low and not math.isclose( + sparse_seconds, candidate_low + ): + winner = transport + elif sparse_seconds > candidate_high and not math.isclose( + sparse_seconds, candidate_high + ): + winner = "nccl" + else: + winner = "depends" + + return Estimate( + transport=transport, + compression=compression, + model_size_gb=model_size_gb, + changed_pct=changed_pct, + sparse_seconds=sparse_seconds, + approximate_wire_gb=( + model_size_gb * changed_pct / 100.0 * _WIRE_MULTIPLIER[compression] + ), + nccl_ib_low_s=nccl_low_s, + nccl_ib_high_s=nccl_high_s, + candidate_ethernet_gbps=candidate_ethernet_gbps, + nccl_ethernet_low_s=candidate_low, + nccl_ethernet_high_s=candidate_high, + break_even_ethernet_low_gbps=break_even_low, + break_even_ethernet_high_gbps=break_even_high, + candidate_winner=winner, + model_size_extrapolated=( + model_size_extrapolated + or not _CALIBRATED_MODEL_SIZE_GB[0] + <= model_size_gb + <= _CALIBRATED_MODEL_SIZE_GB[1] + ), + density_extrapolated=( + not _CALIBRATED_DENSITIES[0] <= changed_pct <= _CALIBRATED_DENSITIES[1] + ), + ) + + +def _positive_float(value: str) -> float: + number = float(value) + if number <= 0: + raise argparse.ArgumentTypeError("must be positive") + return number + + +def parse_args() -> argparse.Namespace: + """Parse command-line arguments.""" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--model-size-gb", type=_positive_float, required=True) + parser.add_argument( + "--changed-pct", + "--sparsity-pct", + type=_positive_float, + required=True, + help="Any positive percentage of changed weight elements.", + ) + parser.add_argument("--transport", choices=("all", "s3", "zmq"), default="all") + parser.add_argument("--compression", choices=("raw", "zstd"), default="zstd") + parser.add_argument( + "--candidate-ethernet-gbps", + type=_positive_float, + help="Raw per-rank Ethernet rate used by the projected NCCL refit.", + ) + parser.add_argument("--json", action="store_true") + return parser.parse_args() + + +def _seconds_range(low: float, high: float) -> str: + return f"{low:.3f}-{high:.3f}s" + + +def _print_results(results: list[Estimate]) -> None: + first = results[0] + print( + f"Model: {first.model_size_gb:g} GB indexed BF16; " + f"changed: {first.changed_pct:g}%; compression: {first.compression}" + ) + print( + "Measured NCCL on 400 Gbps/rank IB: " + f"{_seconds_range(first.nccl_ib_low_s, first.nccl_ib_high_s)}" + ) + + if first.candidate_ethernet_gbps is not None: + assert first.nccl_ethernet_low_s is not None + assert first.nccl_ethernet_high_s is not None + print( + f"Projected NCCL on {first.candidate_ethernet_gbps:g} Gbps/rank " + f"Ethernet: {_seconds_range(first.nccl_ethernet_low_s, first.nccl_ethernet_high_s)}" + ) + + print() + winner_header = "Winner@candidate" if first.candidate_ethernet_gbps else "" + print( + f"{'Path':<6} {'Sparse':>10} {'Wire':>10} " + f"{'Ethernet crossover':>26} {winner_header:>18}" + ) + for result in results: + crossover = ( + f"{result.break_even_ethernet_low_gbps:.2f}-" + f"{result.break_even_ethernet_high_gbps:.2f} Gbps/rank" + ) + winner = result.candidate_winner or "" + print( + f"{result.transport.upper():<6} {result.sparse_seconds:>9.3f}s " + f"{result.approximate_wire_gb:>8.3f} GB {crossover:>26} {winner:>18}" + ) + + print() + print( + "Below the lower crossover, sparse refit beats the full NCCL envelope; " + "above the upper crossover, NCCL wins." + ) + print( + "Projection: T_ethernet = T_H100_IB * 400 / candidate_gbps, with all " + "bandwidth values expressed per rank." + ) + if first.model_size_extrapolated: + print("Note: model size is outside the measured calibration range.") + if first.density_extrapolated: + print("Note: changed density is extrapolated from measured 3% and 5% arms.") + + +def main() -> None: + """Run the estimator.""" + args = parse_args() + transports: tuple[Transport, ...] = ( + ("s3", "zmq") if args.transport == "all" else (args.transport,) + ) + results = [ + estimate( + model_size_gb=args.model_size_gb, + changed_pct=args.changed_pct, + transport=transport, + compression=args.compression, + candidate_ethernet_gbps=args.candidate_ethernet_gbps, + ) + for transport in transports + ] + if args.json: + print(json.dumps([asdict(result) for result in results], indent=2)) + else: + _print_results(results) + + +if __name__ == "__main__": + main() From 4ab3f337bb3190d4dbe70345fa655d516ae5a5cd Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Sat, 11 Jul 2026 12:21:14 -0700 Subject: [PATCH 06/18] Clean up Signed-off-by: Hollow Man --- docs/design-docs/sparse-delta-refit.md | 76 +- nemo_rl/algorithms/grpo.py | 11 +- nemo_rl/models/generation/__init__.py | 5 +- nemo_rl/models/generation/vllm/config.py | 4 +- .../models/generation/vllm/vllm_backend.py | 5 +- .../generation/vllm/vllm_sparse_delta.py | 927 ++++++------------ .../generation/vllm/vllm_sparse_refit.py | 89 +- .../policy/workers/megatron_policy_worker.py | 2 +- .../workers/megatron_remote_sparse_refit.py | 6 +- .../utils/weight_transfer_remote_sparse.py | 39 +- nemo_rl/utils/weight_transfer_sparse_codec.py | 13 +- nemo_rl/utils/weight_transfer_zmq.py | 6 +- .../vllm_remote_sparse_weight_synchronizer.py | 179 ++-- .../generation/test_vllm_sparse_delta.py | 416 ++++---- .../generation/test_vllm_sparse_refit.py | 19 +- .../test_weight_transfer_remote_sparse.py | 128 +-- ..._vllm_remote_sparse_weight_synchronizer.py | 4 +- tools/refit_bandwidth_calculator.py | 228 ++--- 18 files changed, 823 insertions(+), 1334 deletions(-) diff --git a/docs/design-docs/sparse-delta-refit.md b/docs/design-docs/sparse-delta-refit.md index 96d01345388..62932993f01 100644 --- a/docs/design-docs/sparse-delta-refit.md +++ b/docs/design-docs/sparse-delta-refit.md @@ -79,8 +79,9 @@ workers. It uses file-backed `torch.from_file` tensors by default; On a fresh run, vLLM already holds the shared checkpoint. Baseline construction starts early and can overlap initial generation, so the redundant initial full -sync is skipped. A resumed run performs a refit before generation because the -training checkpoint may be newer than the rollout checkpoint. +sync is skipped. On resume, both clusters must still start from the same HF +weight version; sparse refit does not reconstruct a rollout baseline from an +arbitrary training checkpoint. ### Export and encode deltas @@ -98,7 +99,7 @@ expected = previous_baseline + delta.to(baseline_dtype) ``` The producer overlaps export, encoding, `torch.save` serialization, zstd level -1 compression, and transfer with bounded executors. Source baselines do not +1 compression, and transfer with fixed-size executors. Source baselines do not commit until the entire transfer succeeds. ### Transfer and apply @@ -118,7 +119,7 @@ the same receiver endpoints and checksum validation. The receiver deduplicates payload identities, batches them in a bounded FIFO queue, and applies batches on one worker thread. When all vLLM ranks share a node, payloads are staged under `/dev/shm` and passed to collective RPC by file -path. Otherwise, each serialized payload is sent through collective RPC. +path. Otherwise, the serialized batch is sent through one collective RPC. The final `/nemo-rl/refit/flush` drains the queue, synchronizes CUDA, and checks optional delta samples. Only then does the source commit pending baseline @@ -144,15 +145,16 @@ and optional verification samples. HF coordinates are the canonical wire format because Megatron Bridge already defines the training-to-HF mapping while vLLM owns a different packed and -sharded layout. `_SparseDeltaTargetPlan` converts source coordinates to local -vLLM indices without materializing a dense HF tensor. Plans cover identity and -single-dimension shards, packed QKV, merged gate/up projections, fused MoE -experts, and Mamba2 layouts. +sharded layout. On first use, the receiver runs vLLM's native `load_weights()` +against metadata-only tensors while a PyTorch dispatch mode records the source +and destination views of each `copy_`. It caches those mappings and applies +later sparse deltas directly with `index_add_`, without materializing a dense HF +tensor or duplicating QKV, MoE, Mamba, or TP placement rules. -`_SparseDeltaTargetPlan(target=None)` means a valid tensor is absent from the -current rank. A `None` plan means the layout is unsupported; the receiver fails -the payload before applying it. There is no dense fallback for an unknown -layout. +The tracer accepts affine tensor views and the Mamba `A_log` transform. An +element-expanding copy, unknown transform, or unplaced +non-expert tensor fails before any payload update. There is no dense fallback +for an unknown layout. ## Configuration @@ -233,12 +235,13 @@ not duplicate the baseline tracker, codec, receiver queue, or placement logic. Retries must preserve payload identity and bytes, fan out to every required replica, and require a successful global flush before baseline commit. -For a new vLLM layout, derive shard ownership and offsets from vLLM module -attributes or its loader contract. Validate every source shape and target -capacity. Unit tests must exercise nonzero TP ranks, replicated KV heads, local -and remote experts, uneven shapes, contiguous ranges, and explicit locations. -In-range but incorrect `index_add_` locations silently corrupt weights, so test -the exact mapped indices and values. +Do not add model-specific placement math. New layouts should work through their +native vLLM weight loader; extend the tracer only for a general loader operation +and fail closed for transformed or broadcasting copies. Tests must invoke the +real vLLM loader at nonzero TP ranks and cover replicated KV heads, packed +columns, local and remote experts, segmented views, contiguous ranges, and +explicit locations. Incorrect in-range `index_add_` locations silently corrupt +weights, so assert exact mapped indices and values. Codec changes must update encoder and decoder together, preserve 64-bit-safe locations, and retain wire-dtype rounding in pending baseline updates. Receiver @@ -248,13 +251,14 @@ propagation, flush, CUDA synchronization, and clean shutdown. Run the focused suite: ```bash -uv run pytest -q \ +uv run --extra vllm pytest -q \ tests/unit/utils/test_weight_transfer_remote_sparse.py \ tests/unit/models/policy/test_megatron_remote_sparse_refit.py \ tests/unit/models/generation/test_vllm_sparse_refit.py \ - tests/unit/models/generation/test_vllm_sparse_delta.py \ - tests/unit/weight_sync/test_vllm_remote_sparse_weight_synchronizer.py \ - tests/unit/tools/test_refit_bandwidth_calculator.py + tests/unit/weight_sync/test_vllm_remote_sparse_weight_synchronizer.py + +uv run --extra vllm pytest -q -m vllm \ + tests/unit/models/generation/test_vllm_sparse_delta.py uv run ruff check \ nemo_rl/utils/weight_transfer_{remote_sparse,sparse_codec,zmq}.py \ @@ -264,7 +268,7 @@ uv run ruff check \ ``` On the target topology, verify the exact commit, image digest, and checkpoint -revision; run fresh and resumed starts; compare at least two balanced +revision; validate fresh starts and same-version resumes; compare two balanced repetitions with an equivalent NCCL or full control; and require the requested changed density, one global commit, no traceback, and zero sampled mismatches. After failure injection, confirm the source baseline does not commit and reload @@ -276,23 +280,15 @@ the receiver before retrying. benchmark-specific estimator for the current S3 and ZeroMQ implementation. It is not a general fabric or topology model. -The sparse side embeds July 2026 end-to-end fits from 32 GB300 sender GPUs in -`us-east-2` to 64 H100 receiver GPUs in `us-east-1`. The measured checkpoints -span 63.2-1121.0 GB of indexed BF16 weights. With `S` as decimal TB, the -embedded latency fits are: - -| Transport | Payload | 3% changed | 5% changed | -|---|---|---:|---:| -| S3 | raw | `2.370 + 116.138S` | `2.646 + 183.431S` | -| S3 | zstd | `84.746S` | `134.041S` | -| ZeroMQ | raw | `264.753S` | `428.008S` | -| ZeroMQ | zstd | `6.517 + 73.621S` | `8.165 + 157.339S` | - -For an arbitrary positive changed density `d`, the calculator evaluates `T3` -and `T5` from this table, computes `p = log(T5 / T3) / log(5 / 3)`, and returns -`T3 * (d / 3)^p`. Inputs outside 3-5% are accepted but are extrapolations. The -wire estimate is `model_size_gb * d / 100 * multiplier`, where the raw and zstd -multipliers are 2.0 and 0.74. +The sparse side embeds July 2026 end-to-end latency fits from 32 GB300 sender +GPUs in `us-east-2` to 64 H100 receiver GPUs in `us-east-1`. Measurements cover +63.2-1121.0 GB of indexed BF16 weights, S3 and ZeroMQ, raw and zstd payloads, +and 3% and 5% changed density. Any positive `--changed-pct` is accepted; values +outside 3-5% use an extrapolated power curve through the two measured-density +fits. Estimated wire bytes use the measured raw or zstd payload ratio. +The coefficients in `_SPARSE_LATENCY_FITS` implement +`fixed_seconds + seconds_per_1000_GB * model_size_GB / 1000`; they are latency +regressions, not bandwidth measurements. The NCCL side uses these measured generation-EP H100 refit envelopes on 400 Gbps/rank InfiniBand: diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index 0599d2e2662..b09795f73a2 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -2079,13 +2079,7 @@ def refit_policy_generation( """ synchronizer = getattr(policy_generation, "weight_synchronizer", None) if synchronizer is not None: - return ( - synchronizer.sync_weights( - timer=timer, - kv_scales=kv_scales, - ) - or {} - ) + return synchronizer.sync_weights(timer=timer, kv_scales=kv_scales) or {} # Megatron generation backend needs explicit suspend/resume around refits. if isinstance(policy_generation, MegatronGeneration): @@ -2465,10 +2459,10 @@ def grpo_train( batch_cache: BatchedDataDict[DatumSpec] = None # This is the number of batches we processed so far at each step to generate responses whose std is non-zero. Maximum threshold is set by dynamic_sampling_max_gen_batches. Used in the case of dynamic sampling. dynamic_sampling_num_gen_batches = 0 - refit_metrics: dict[str, float] = {} # Run grpo/dapo training loop (single-turn) for batch in wrapped_dataloader: + refit_metrics: dict[str, float] = {} # A central place to store logging data that won't be deleted until the loop ends metrics_logging_data = dict() metrics = dict() @@ -3369,7 +3363,6 @@ def grpo_train( # Reset the batch and set dynamic_sampling_num_gen_batches to 0 batch_cache = None dynamic_sampling_num_gen_batches = 0 - refit_metrics = {} # Clear mem memory_tracker.snapshot_start_of_stage("After CPU memory clear", dir()) diff --git a/nemo_rl/models/generation/__init__.py b/nemo_rl/models/generation/__init__.py index f208a40a945..cc471b1b27b 100644 --- a/nemo_rl/models/generation/__init__.py +++ b/nemo_rl/models/generation/__init__.py @@ -46,10 +46,7 @@ def configure_generation_config( config = cast(VllmConfig, config) # set load_format config["vllm_cfg"]["load_format"] = ( - "auto" - if is_eval - or config.get("refit_transport") in ("vllm_s3_sparse", "vllm_zmq_sparse") - else "dummy" + "auto" if is_eval or config.get("refit_transport") else "dummy" ) speculative_config = config.get("vllm_kwargs", {}).get("speculative_config") if speculative_config and not is_eval and not has_refit_draft_weights: diff --git a/nemo_rl/models/generation/vllm/config.py b/nemo_rl/models/generation/vllm/config.py index 0f4a86b687e..62329563b34 100644 --- a/nemo_rl/models/generation/vllm/config.py +++ b/nemo_rl/models/generation/vllm/config.py @@ -16,8 +16,6 @@ from nemo_rl.models.generation.interfaces import GenerationConfig -DeltaCompressionDType = Literal["fp16", "float16", "bf16", "bfloat16", "fp32", "float32"] # fmt: skip - class VllmSpecificArgs(TypedDict): tensor_parallel_size: int @@ -61,7 +59,7 @@ class VllmSpecificArgs(TypedDict): class VllmDeltaCompressionConfig(TypedDict): - dtype: DeltaCompressionDType + dtype: Literal["fp16", "float16", "bf16", "bfloat16", "fp32", "float32"] # fmt: skip sparse_bucket_size_bytes: int diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index 5cadc893252..c641fa74471 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -372,7 +372,6 @@ def _get_sparse_delta_applier(self) -> Any: self._sparse_delta_applier = VllmSparseDeltaApplier( self.model_runner, self.device, - rank=int(getattr(self, "rank", 0)), ) return self._sparse_delta_applier @@ -507,10 +506,10 @@ def update_weights_from_collective(self) -> bool: def update_weights_from_serialized_sparse_payload( self, - serialized_payload: bytes, + *serialized_payloads: bytes, ) -> dict[str, Any]: return self._get_sparse_delta_applier().update_weights_from_serialized_sparse_payload( - serialized_payload + *serialized_payloads ) def update_weights_from_sparse_payload_files( diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py index d26497aa1d6..e0569c8b6fc 100644 --- a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py +++ b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py @@ -17,707 +17,366 @@ import io import re import time +from collections import defaultdict from dataclasses import dataclass +from math import prod from typing import Any, cast import torch +from torch.utils._python_dispatch import TorchDispatchMode from nemo_rl.utils import weight_transfer_sparse_codec as sparse_codec from nemo_rl.utils.nsys import wrap_with_nvtx_name -_EXPERT_WEIGHT_RE = re.compile( - r"^(?P.*\.experts)\.(?P\d+)\." - r"(?Pgate_proj|up_proj|down_proj)\.weight$" -) +_EXPERT_WEIGHT_RE = re.compile(r"\.experts\.\d+\.(?:gate|up|down)_proj\.weight$") + + +def _storage_key(tensor: torch.Tensor) -> int: + return tensor.untyped_storage()._cdata + + +@dataclass(frozen=True) +class _SparseDeltaCopyPlan: + target: torch.Tensor + source_offset: int + source_strides: tuple[int, ...] + shape: tuple[int, ...] + target_offset: int + target_strides: tuple[int, ...] + linear: bool @dataclass(frozen=True) class _SparseDeltaTargetPlan: - target: torch.Tensor | None - source_shape: tuple[int, ...] = () - source_strides: tuple[int, ...] = () - target_strides: tuple[int, ...] = () - target_offset: int = 0 - shard_dim: int | None = None - shard_start: int = 0 - shard_size: int = 0 - segment_shards: tuple[tuple[int, int, int], ...] = () + copies: tuple[_SparseDeltaCopyPlan, ...] = () log_delta_transform: bool = False identity: bool = False -class VllmSparseDeltaApplier: - """Own sparse placement state without extending the normal refit path.""" +class _SparseLoadTracer(TorchDispatchMode): + """Capture the views copied by vLLM's native weight loaders.""" def __init__( self, - model_runner: Any, - device: torch.device, - *, - rank: int = 0, + targets: list[torch.Tensor], + sources: dict[str, torch.Tensor], ) -> None: + super().__init__() + self.copies: dict[str, list[_SparseDeltaCopyPlan]] = defaultdict(list) + self.postprocessed: set[str] = set() + self._sources = {_storage_key(tensor): name for name, tensor in sources.items()} + self._targets: dict[int, list[torch.Tensor]] = defaultdict(list) + for target in targets: + if not target.is_contiguous(): + raise RuntimeError("Sparse delta targets must be contiguous.") + self._targets[_storage_key(target)].append(target) + self._last_source: dict[int, str] = {} + + def _target_for(self, view: torch.Tensor) -> torch.Tensor | None: + candidates = self._targets.get(_storage_key(view), ()) + view_start = view.storage_offset() + view_end = view_start + sum( + (size - 1) * stride for size, stride in zip(view.shape, view.stride()) + ) + for target in candidates: + start = target.storage_offset() + if start <= view_start and view_end < start + target.numel(): + return target + return None + + def __torch_dispatch__( + self, + func: Any, + types: Any, + args: tuple[Any, ...] = (), + kwargs: dict[str, Any] | None = None, + ) -> Any: + if func is not torch.ops.aten.copy_.default: + return func(*args, **(kwargs or {})) + + destination, source = cast(tuple[torch.Tensor, torch.Tensor], args[:2]) + target = self._target_for(destination) + if target is None: + raise RuntimeError("vLLM loader copied outside a model parameter.") + + source_name = self._sources.get(_storage_key(source)) + target_key = id(target) + if source_name is None: + source_name = self._last_source.get(target_key) + if source_name is None: + raise RuntimeError("vLLM loader materialized an unsupported transform.") + self.postprocessed.add(source_name) + return destination + + if destination.shape != source.shape: + raise RuntimeError("vLLM loader used an expanding copy.") + shape = tuple(source.shape) + source_strides = tuple(source.stride()) + target_strides = tuple(destination.stride()) + if any(size > 1 and stride <= 0 for size, stride in zip(shape, source_strides)): + raise RuntimeError("vLLM loader used an unsupported source view.") + contiguous = torch.empty(shape, device="meta").stride() + self.copies[source_name].append( + _SparseDeltaCopyPlan( + target, + int(source.storage_offset()), + source_strides, + shape, + int(destination.storage_offset() - target.storage_offset()), + target_strides, + source_strides == target_strides == contiguous, + ) + ) + self._last_source[target_key] = source_name + return destination + + +class VllmSparseDeltaApplier: + """Apply sparse HF deltas through plans derived from native vLLM loaders.""" + + def __init__(self, model_runner: Any, device: torch.device) -> None: self.model_runner = model_runner self._cuda_device_index = device.index - self.rank = rank - self._direct_sparse_delta_targets: dict[str, torch.Tensor] | None = None - self._direct_sparse_delta_plan_cache: dict[ - str, _SparseDeltaTargetPlan | None - ] = {} - self._direct_sparse_delta_verification: list[ + self._plan_cache: dict[str, _SparseDeltaTargetPlan] = {} + self._verification: list[ tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] ] = [] - self._direct_sparse_delta_verification_candidates = 0 + self._verification_candidates = 0 - def _apply_sparse_weight_deltas( - self, - payload_tensors: tuple[torch.Tensor, torch.Tensor], - metadata: list[dict[str, Any]], - ) -> None: - """Apply sparse deltas directly after validating every target plan.""" - architectures = self.model_runner.vllm_config.model_config.architectures - # Delay the vLLM-dependent FP8 helper until a payload is applied. - from nemo_rl.models.generation.vllm.quantization import fp8 - - if {"GptOssForCausalLM", "Gemma3ForConditionalGeneration"} & set( - architectures - ) or fp8.is_fp8_model(self.model_runner.vllm_config): - raise RuntimeError( - "Direct sparse delta refit does not support transformed or FP8 weights." - ) + def _compile_plans(self, metadata: list[dict[str, Any]]) -> None: + missing = { + str(item["name"]): item + for item in metadata + if item["name"] not in self._plan_cache + } + if not missing: + return + + model = self.model_runner.model + targets = list(model.parameters()) + list(model.buffers()) + sources = { + name: torch.empty(tuple(item["shape"]), device="meta") + for name, item in missing.items() + } + tracer = _SparseLoadTracer(targets, sources) + with torch.no_grad(), tracer: + model.load_weights((name, sources[name]) for name in missing) - if self._direct_sparse_delta_targets is None: - model = self.model_runner.model - self._direct_sparse_delta_targets = dict(model.named_parameters()) | dict( - model.named_buffers() + for name, item in missing.items(): + source_shape = tuple(item["shape"]) + source_strides = tuple(sources[name].stride()) + copies = tuple(tracer.copies.get(name, ())) + transformed = name in tracer.postprocessed + log_transform = ( + transformed and ".mixer." in name and name.endswith((".A", ".A_log")) ) - targets = self._direct_sparse_delta_targets - raw_locations, raw_values = payload_tensors - plan_cache = self._direct_sparse_delta_plan_cache - if plan_cache is None: - plan_cache = self._direct_sparse_delta_plan_cache = {} - plans = [] - for item in metadata: - name = str(item["name"]) - if name not in plan_cache: - plan_cache[name] = self._direct_sparse_delta_target_plan(item, targets) - plan = plan_cache[name] - if plan is None: + if transformed and not log_transform: raise RuntimeError( - f"No direct sparse delta target plan for {item['name']!r}." + f"vLLM loader for {name!r} transforms weights and cannot apply deltas." ) - plans.append((item, plan)) - - with torch.no_grad(): - for item, plan in plans: - target = plan.target - verification_locations = item.get("verification_locations", []) - self._direct_sparse_delta_verification_candidates += len( - verification_locations - ) - if target is None: - continue - - if verification_locations and not plan.log_delta_transform: - sample_locations, sample_deltas = ( - self._local_sparse_delta_update_inputs( - torch.tensor(verification_locations, device=target.device), - torch.tensor( - item["verification_deltas"], - device=target.device, - dtype=target.dtype, - ), - plan, - ) - ) - if sample_locations.numel(): - before = target.data.view(-1).index_select(0, sample_locations) - expected_delta = ( - before + sample_deltas - ).float() - before.float() - verification = self._direct_sparse_delta_verification - if verification is None: - verification = self._direct_sparse_delta_verification = [] - verification.append( - ( - target, - sample_locations, - before.float(), - expected_delta, - ) - ) - - value_start = int(item["value_start"]) - value_end = int(item["value_end"]) - values = raw_values[value_start:value_end].to( - device=target.device, - dtype=target.dtype, - non_blocking=True, - ) - if plan.identity and item["index_encoding"] == "range": - range_start = int(item["range_start"]) - range_count = value_end - value_start - target.data.view(-1).narrow(0, range_start, range_count).add_( - values - ) - else: - locations = sparse_codec.sparse_locations_for_item( - item, - raw_locations, - device=target.device, - ) - locations, values = self._local_sparse_delta_update_inputs( - locations, - values, - plan, - ) - if locations.numel(): - target_flat = target.data.view(-1) - if plan.log_delta_transform: - current = target_flat.index_select(0, locations) - updated = current * values.float().exp().to( - dtype=current.dtype - ) - target_flat.index_copy_(0, locations, updated) - else: - target_flat.index_add_(0, locations, values) - - def _direct_sparse_delta_module( - self, - target: torch.Tensor, - module_name: str, - ) -> Any: - loader = getattr(target, "weight_loader", None) - return getattr( - loader, "__self__", None - ) or self.model_runner.model.get_submodule(module_name) - - def _direct_sparse_delta_target_plan( - self, - item: dict[str, Any], - targets: dict[str, torch.Tensor], - ) -> _SparseDeltaTargetPlan | None: - name = str(item["name"]) - if name.startswith("mtp."): - return _SparseDeltaTargetPlan(target=None) - mapper = getattr(self.model_runner.model, "hf_to_vllm_mapper", None) - target_name = cast(Any, mapper)._map_name(name) if mapper is not None else name - if target_name is None or target_name.startswith("draft."): - return None - if ".mixer." in target_name: - mamba_plan = self._direct_sparse_delta_mamba2_plan( - item, target_name, targets + if not copies and not ( + name.startswith(("mtp.", "draft.")) or _EXPERT_WEIGHT_RE.search(name) + ): + raise RuntimeError(f"vLLM loader did not place {name!r}.") + identity = ( + len(copies) == 1 + and not log_transform + and copies[0].source_offset == 0 + and copies[0].shape == source_shape + and copies[0].source_strides == source_strides + and copies[0].target_offset == 0 + and copies[0].target_strides == source_strides + and copies[0].target.numel() == prod(source_shape) ) - if mamba_plan is not None: - return mamba_plan - if ".mixer.conv1d." in target_name or ".mixer.in_proj." in target_name: - return None - if any(f".{candidate}_proj." in target_name for candidate in ("q", "k", "v")): - return self._direct_sparse_delta_qkv_plan(item, target_name, targets) - if _EXPERT_WEIGHT_RE.match(target_name): - return self._direct_sparse_delta_expert_plan(item, target_name, targets) - if any(f".{candidate}_proj." in target_name for candidate in ("gate", "up")): - merged_plan = self._direct_sparse_delta_merged_column_plan( - item, target_name, targets + self._plan_cache[name] = _SparseDeltaTargetPlan( + copies, log_transform, identity ) - if merged_plan is not None: - return merged_plan - target = targets.get(target_name) - if target is None: - return None - source_shape = tuple(item["shape"]) - target_shape = tuple(target.shape) - if target_shape == source_shape: - return self._make_sparse_delta_target_plan(target, source_shape) - return self._direct_sparse_delta_shard_plan(item, target) - - def _direct_sparse_delta_qkv_plan( - self, - item: dict[str, Any], - target_name: str, - targets: dict[str, torch.Tensor], - ) -> _SparseDeltaTargetPlan | None: - shard_id = next(x for x in "qkv" if f".{x}_proj." in target_name) - packed_name = target_name.replace(f".{shard_id}_proj.", ".qkv_proj.", 1) - target = targets.get(packed_name) - if target is None: - return None - output_dim = int(cast(Any, target).output_dim) % target.ndim - module = self._direct_sparse_delta_module(target, packed_name.rsplit(".", 1)[0]) - shard_offset = int(module._get_shard_offset_mapping(shard_id)) - shard_size = int(module._get_shard_size_mapping(shard_id)) - shard_rank = int(module.tp_rank) - if shard_id != "q": - shard_rank //= int(module.num_kv_head_replicas) - - source_shape = tuple(item["shape"]) - shard_start = shard_rank * shard_size - if source_shape[output_dim] < shard_start: - return _SparseDeltaTargetPlan(target=None) - return self._make_sparse_delta_target_plan( - target, - source_shape=source_shape, - shard_dim=output_dim, - shard_start=shard_start, - shard_size=min(shard_size, source_shape[output_dim] - shard_start), - target_offset=shard_offset * target.stride(output_dim), - ) - - def _direct_sparse_delta_merged_column_plan( - self, - item: dict[str, Any], - target_name: str, - targets: dict[str, torch.Tensor], - ) -> _SparseDeltaTargetPlan | None: - projection = next( - candidate - for candidate in ("gate", "up") - if f".{candidate}_proj." in target_name - ) - shard_id = 0 if projection == "gate" else 1 - packed_name = target_name.replace(f".{projection}_proj.", ".gate_up_proj.", 1) - target = targets.get(packed_name) - output_dim = getattr(target, "output_dim", None) - if target is None or not isinstance(output_dim, int): - return None - - output_dim %= target.ndim - module = self._direct_sparse_delta_module(target, packed_name.rsplit(".", 1)[0]) - output_sizes = tuple(int(size) for size in module.output_sizes) - tp_size = int(module.tp_size) - source_shape = tuple(item["shape"]) - if ( - shard_id >= len(output_sizes) - or tp_size < 1 - or output_sizes[shard_id] % tp_size - or output_dim >= len(source_shape) - or source_shape[output_dim] != output_sizes[shard_id] + @staticmethod + def _map_copy( + locations: torch.Tensor, + values: torch.Tensor, + copy: _SparseDeltaCopyPlan, + ) -> tuple[torch.Tensor, torch.Tensor]: + if copy.linear: + end = copy.source_offset + prod(copy.shape) + keep = (locations >= copy.source_offset) & (locations < end) + return ( + locations[keep] + copy.target_offset - copy.source_offset, + values[keep], + ) + mapped = torch.full_like(locations, copy.target_offset) + reconstructed = torch.full_like(locations, copy.source_offset) + relative = locations - copy.source_offset + for size, source_stride, target_stride in zip( + copy.shape, copy.source_strides, copy.target_strides, strict=True ): - return None - - shard_size = output_sizes[shard_id] // tp_size - target_start = sum(output_sizes[:shard_id]) // tp_size - if target.shape[output_dim] < target_start + shard_size: - return None - shard_start = int(module.tp_rank) * shard_size - return self._make_sparse_delta_target_plan( - target, - source_shape=source_shape, - shard_dim=output_dim, - shard_start=shard_start, - shard_size=shard_size, - target_offset=target_start * target.stride(output_dim), - ) - - def _direct_sparse_delta_mamba2_plan( - self, - item: dict[str, Any], - target_name: str, - targets: dict[str, torch.Tensor], - ) -> _SparseDeltaTargetPlan | None: - target = targets.get(target_name) - if target is None: - return None - - if target_name.endswith(".A"): - source_shape = tuple(item["shape"]) - if tuple(target.shape) == source_shape: - return self._make_sparse_delta_target_plan( - target, - source_shape=source_shape, - log_delta_transform=True, + coordinate = ( + torch.div(relative, source_stride, rounding_mode="floor").remainder( + size ) - return self._direct_sparse_delta_shard_plan( - item, - target, - log_delta_transform=True, + if size > 1 + else torch.zeros_like(locations) ) + reconstructed.add_(coordinate * source_stride) + mapped.add_(coordinate * target_stride) + keep = reconstructed == locations + return mapped[keep], values[keep] - if not (".mixer.conv1d." in target_name or ".mixer.in_proj." in target_name): - return None - - mixer_name = target_name.split(".mixer.", 1)[0] + ".mixer" - attrs = cast(Any, self.model_runner.model.get_submodule(mixer_name)) - tp_size = int(attrs.tp_size) - if tp_size <= 1: - return None - intermediate_size = int(attrs.intermediate_size) - groups_ssm_state_size = int(attrs.groups_ssm_state_size) - num_heads = int(attrs.num_heads) - source_shape = tuple(item["shape"]) - fixed_size = intermediate_size - if ".mixer.in_proj." in target_name: - fixed_size += intermediate_size + num_heads - group_size, remainder = divmod(source_shape[0] - fixed_size, 2) - extra_group_size = groups_ssm_state_size - group_size - if remainder or group_size <= 0 or extra_group_size < 0: - return None - tp_rank = int( - getattr(target, "tp_rank", int(getattr(self, "rank", 0)) % tp_size) - ) - intermediate = (intermediate_size, 0, False) - group = (groups_ssm_state_size, extra_group_size, extra_group_size > 0) - segment_specs = ( - (intermediate, group, group) - if ".mixer.conv1d." in target_name - else (intermediate, intermediate, group, group, (num_heads, 0, False)) - ) - - target_shape = tuple(target.shape) - source_to_target_dims = tuple(range(len(source_shape))) - if len(target_shape) == len(source_shape) + 1 and target_shape[1] == 1: - source_to_target_dims = (0, *range(2, len(target_shape))) - elif len(target_shape) != len(source_shape): - return None - segment_shards: list[tuple[int, int, int]] = [] - target_start = 0 - source_start = 0 - for full_dim, extra, duplicate_groups in segment_specs: - shard_size = full_dim // tp_size - rank = 0 if duplicate_groups else tp_rank - source_dim = full_dim - extra - source_local_start = source_start + rank * shard_size - take = min(shard_size, source_dim - rank * shard_size) - if take > 0: - segment_shards.append((source_local_start, target_start, take)) - target_start += shard_size - source_start += source_dim - if source_shape[0] != source_start or target_shape[0] != target_start: - return None - - return self._make_sparse_delta_target_plan( - target, - source_shape=source_shape, - source_to_target_dims=source_to_target_dims, - shard_dim=0, - segment_shards=tuple(segment_shards), - ) - - def _direct_sparse_delta_expert_plan( + def _record_verification( self, item: dict[str, Any], - target_name: str, - targets: dict[str, torch.Tensor], - ) -> _SparseDeltaTargetPlan | None: - match = cast(re.Match[str], _EXPERT_WEIGHT_RE.match(target_name)) - - prefix = match.group("prefix") - global_expert_id = int(match.group("expert")) - proj = match.group("proj") - packed_weight, shard_id = { - "gate_proj": ("w13_weight", "w1"), - "up_proj": ("w13_weight", "w3"), - "down_proj": ("w2_weight", "w2"), - }[proj] - packed_name = f"{prefix}.{packed_weight}" - - target = targets.get(packed_name) - if target is None: - return None - module_attrs = self._direct_sparse_delta_module( - target, packed_name.rsplit(".", 1)[0] - ) - if shard_id == "w3" and not module_attrs.moe_config.is_act_and_mul: - shard_id = "w1" - local_expert_id = int( - module_attrs._map_global_expert_id_to_local_expert_id(global_expert_id) - ) - if local_expert_id < 0: - return _SparseDeltaTargetPlan(target=None) - - source_shape = tuple(item["shape"]) - target_shape = tuple(target.shape) - shard_dim = 1 if shard_id == "w2" else 0 - if local_expert_id >= target_shape[0]: - return None - target_shard_dim = shard_dim + 1 - - shard_size = target_shape[target_shard_dim] - if shard_id in ("w1", "w3") and module_attrs.moe_config.is_act_and_mul: - shard_size //= 2 - target_shard_offset = shard_size if shard_id == "w3" else 0 - if target_shape[target_shard_dim] < target_shard_offset + shard_size: - return None - tp_rank = int(module_attrs.tp_rank) - shard_start = tp_rank * shard_size - if source_shape[shard_dim] < shard_start: - return _SparseDeltaTargetPlan(target=None) - - return self._make_sparse_delta_target_plan( - target, - source_shape=source_shape, - source_to_target_dims=tuple(dim + 1 for dim in range(len(source_shape))), - target_offset=( - local_expert_id * target.stride(0) - + target_shard_offset * target.stride(target_shard_dim) - ), - shard_dim=shard_dim, - shard_start=shard_start, - shard_size=min(shard_size, source_shape[shard_dim] - shard_start), - ) - - def _make_sparse_delta_target_plan( - self, - target: torch.Tensor, - source_shape: tuple[int, ...], - *, - source_to_target_dims: tuple[int, ...] | None = None, - target_offset: int = 0, - shard_dim: int | None = None, - shard_start: int = 0, - shard_size: int = 0, - segment_shards: tuple[tuple[int, int, int], ...] = (), - log_delta_transform: bool = False, - ) -> _SparseDeltaTargetPlan | None: - if source_to_target_dims is None: - source_to_target_dims = tuple(range(len(source_shape))) - target_shape = tuple(target.shape) - ignored_dim = ( - shard_dim if shard_dim is not None else 0 if segment_shards else -1 - ) - if len(source_to_target_dims) != len(source_shape) or any( - target_dim >= len(target_shape) - or ( - source_dim != ignored_dim - and source_shape[source_dim] != target_shape[target_dim] - ) - for source_dim, target_dim in enumerate(source_to_target_dims) - ): - return None - identity = ( - shard_dim is None - and target_offset == 0 - and not segment_shards - and not log_delta_transform - and source_to_target_dims == tuple(range(len(source_shape))) - and source_shape == target_shape - ) - return _SparseDeltaTargetPlan( - target=target, - source_shape=source_shape, - source_strides=torch.empty(source_shape, device="meta").stride(), - target_strides=tuple( - target.stride(target_dim) for target_dim in source_to_target_dims - ), - target_offset=target_offset, - shard_dim=shard_dim, - shard_start=shard_start, - shard_size=shard_size, - segment_shards=segment_shards, - log_delta_transform=log_delta_transform, - identity=identity, + plan: _SparseDeltaTargetPlan, + ) -> None: + sample_locations = item.get("verification_locations", []) + self._verification_candidates += len(sample_locations) + if not sample_locations or plan.log_delta_transform or not plan.copies: + return + target = plan.copies[0].target + locations = torch.tensor(sample_locations, device=target.device) + values = torch.tensor( + item["verification_deltas"], device=target.device, dtype=target.dtype ) + for copy in plan.copies: + mapped, selected = self._map_copy(locations, values, copy) + if not mapped.numel(): + continue + before = copy.target.data.view(-1).index_select(0, mapped) + expected = (before + selected).float() - before.float() + self._verification.append((copy.target, mapped, before.float(), expected)) - def _direct_sparse_delta_shard_plan( + def _apply_item( self, item: dict[str, Any], - target: torch.Tensor, - *, - log_delta_transform: bool = False, - ) -> _SparseDeltaTargetPlan | None: - source_shape = tuple(item["shape"]) - target_shape = tuple(target.shape) - if len(source_shape) != len(target_shape): - return None - - candidate_dims = list( - dict.fromkeys( - dim % len(source_shape) - for attr in ("output_dim", "input_dim") - if isinstance(dim := getattr(target, attr, None), int) - ) + plan: _SparseDeltaTargetPlan, + raw_locations: torch.Tensor, + raw_values: torch.Tensor, + ) -> None: + self._record_verification(item, plan) + if not plan.copies: + return + first_target = plan.copies[0].target + value_start, value_end = int(item["value_start"]), int(item["value_end"]) + values = raw_values[value_start:value_end].to( + device=first_target.device, dtype=first_target.dtype, non_blocking=True + ) + if plan.identity and item["index_encoding"] == "range": + first_target.data.view(-1).narrow( + 0, int(item["range_start"]), value_end - value_start + ).add_(values) + return + + locations = sparse_codec.sparse_locations_for_item( + item, raw_locations, device=first_target.device ) - if not candidate_dims: - candidate_dims = [ - dim - for dim, (source_dim, target_dim) in enumerate( - zip(source_shape, target_shape, strict=True) + if plan.identity: + first_target.data.view(-1).index_add_(0, locations, values) + return + grouped: dict[ + int, tuple[torch.Tensor, list[torch.Tensor], list[torch.Tensor]] + ] = {} + for copy in plan.copies: + mapped, selected = self._map_copy(locations, values, copy) + if mapped.numel(): + _, mapped_parts, value_parts = grouped.setdefault( + id(copy.target), (copy.target, [], []) ) - if source_dim != target_dim - ] - if len(candidate_dims) != 1: - return None - - for shard_dim in candidate_dims: - shard_size = target_shape[shard_dim] - tp_size = int(getattr(target, "tp_size", 1)) - if tp_size <= 1: - if shard_size <= 0 or source_shape[shard_dim] % shard_size: - continue - tp_size = source_shape[shard_dim] // shard_size - if tp_size <= 1: - continue - if source_shape[shard_dim] > shard_size * tp_size: - continue - tp_rank = int( - getattr(target, "tp_rank", int(getattr(self, "rank", 0)) % tp_size) + mapped_parts.append(mapped) + value_parts.append(selected) + for target, mapped_parts, value_parts in grouped.values(): + mapped = ( + torch.cat(mapped_parts) if len(mapped_parts) > 1 else mapped_parts[0] ) - plan = self._make_sparse_delta_target_plan( - target=target, - source_shape=source_shape, - shard_dim=shard_dim, - shard_start=tp_rank * shard_size, - shard_size=shard_size, - log_delta_transform=log_delta_transform, + selected = ( + torch.cat(value_parts) if len(value_parts) > 1 else value_parts[0] ) - if plan is not None: - return plan - return None + target_flat = target.data.view(-1) + if plan.log_delta_transform: + current = target_flat.index_select(0, mapped) + target_flat.index_copy_( + 0, mapped, current * selected.float().exp().to(current.dtype) + ) + else: + target_flat.index_add_(0, mapped, selected) - def _local_sparse_delta_update_inputs( + def _apply_sparse_weight_deltas( self, - locations: torch.Tensor, - values: torch.Tensor, - plan: _SparseDeltaTargetPlan, - ) -> tuple[torch.Tensor, torch.Tensor]: - if plan.identity: - return locations, values - - source_shape = plan.source_shape - source_strides = plan.source_strides - target_strides = plan.target_strides - shard_dim = plan.shard_dim - - if source_strides == target_strides: - if shard_dim is None: - return locations + plan.target_offset, values - if shard_dim == 0: - shard_stride = source_strides[0] - shard_coords = torch.div(locations, shard_stride, rounding_mode="floor") - if plan.segment_shards: - mapped_locations = locations + plan.target_offset - keep = torch.zeros_like(locations, dtype=torch.bool) - for source_start, target_start, take in plan.segment_shards: - segment = (shard_coords >= source_start) & ( - shard_coords < source_start + take - ) - mapped_locations[segment] += ( - target_start - source_start - ) * shard_stride - keep |= segment - return mapped_locations[keep], values[keep] - shard_end = min( - plan.shard_start + plan.shard_size, source_shape[shard_dim] - ) - keep = (shard_coords >= plan.shard_start) & (shard_coords < shard_end) - return ( - locations[keep] - + plan.target_offset - - plan.shard_start * shard_stride, - values[keep], - ) + payload_tensors: tuple[torch.Tensor, torch.Tensor], + metadata: list[dict[str, Any]], + ) -> None: + from nemo_rl.models.generation.vllm.quantization import fp8 - selected_locations = locations - selected_values = values - - if shard_dim is not None: - shard_coords = torch.div( - locations, - source_strides[shard_dim], - rounding_mode="floor", - ).remainder(source_shape[shard_dim]) - shard_end = min( - plan.shard_start + plan.shard_size, - source_shape[shard_dim], + if fp8.is_fp8_model(self.model_runner.vllm_config): + raise RuntimeError( + "Direct sparse delta refit does not support FP8 weights." ) - keep = (shard_coords >= plan.shard_start) & (shard_coords < shard_end) - selected_locations = locations[keep] - selected_values = values[keep] - if selected_locations.numel() == 0: - return selected_locations, selected_values - - local_locations = torch.full_like(selected_locations, plan.target_offset) - for dim, (source_stride, target_stride) in enumerate( - zip(source_strides, target_strides, strict=True) - ): - coord = torch.div( - selected_locations, - source_stride, - rounding_mode="floor", - ).remainder(source_shape[dim]) - if dim == plan.shard_dim: - coord = coord - plan.shard_start - local_locations.add_(coord * target_stride) - return local_locations, selected_values + + self._compile_plans(metadata) + raw_locations, raw_values = payload_tensors + with torch.no_grad(): + for item in metadata: + self._apply_item( + item, self._plan_cache[str(item["name"])], raw_locations, raw_values + ) @wrap_with_nvtx_name( "vllm_internal_worker_extension/update_weights_from_serialized_sparse_payload" ) def update_weights_from_serialized_sparse_payload( - self, - serialized_payload: bytes, - ) -> dict[str, Any]: - """Apply one serialized sparse-delta payload.""" - return self._load_and_apply_sparse_payload(io.BytesIO(serialized_payload)) - - def _load_and_apply_sparse_payload( - self, - source: str | io.BytesIO, + self, *serialized_payloads: bytes ) -> dict[str, Any]: - started = time.perf_counter() - payload = cast( - sparse_codec.TensorPayload, - torch.load( - source, - map_location="cpu", - weights_only=True, - ), + return self._load_and_apply_sparse_payloads( + tuple(io.BytesIO(payload) for payload in serialized_payloads) ) - deserialize_s = time.perf_counter() - started - result = self._apply_sparse_request(payload) - result["receiver_deserialize_s"] = deserialize_s - result["receiver_total_s"] = time.perf_counter() - started - return result - @wrap_with_nvtx_name( - "vllm_internal_worker_extension/update_weights_from_sparse_payload_files" - ) - def update_weights_from_sparse_payload_files( - self, - *payload_paths: str, + def _load_and_apply_sparse_payloads( + self, sources: tuple[str | io.BytesIO, ...] ) -> dict[str, Any]: - """Apply sparse payloads in FIFO order.""" started = time.perf_counter() - deserialize_s = 0.0 - sparse_apply_s = 0.0 - for path in payload_paths: - result = self._load_and_apply_sparse_payload(path) - deserialize_s += float(result["receiver_deserialize_s"]) - sparse_apply_s += float(result["receiver_sparse_apply_s"]) + deserialize_s = sparse_apply_s = 0.0 + payloads = [] + for source in sources: + item_started = time.perf_counter() + payloads.append( + cast( + sparse_codec.TensorPayload, + torch.load(source, map_location="cpu", weights_only=True), + ) + ) + deserialize_s += time.perf_counter() - item_started + + item_started = time.perf_counter() + self._compile_plans([item for _, _, metadata in payloads for item in metadata]) + plan_s = time.perf_counter() - item_started + for locations, values, metadata in payloads: + item_started = time.perf_counter() + self._apply_sparse_weight_deltas((locations, values), metadata) + sparse_apply_s += time.perf_counter() - item_started return { "ok": True, "receiver_deserialize_s": deserialize_s, + "receiver_plan_s": plan_s, "receiver_sparse_apply_s": sparse_apply_s, "receiver_total_s": time.perf_counter() - started, } - def _apply_sparse_request( - self, - payload: sparse_codec.TensorPayload, + @wrap_with_nvtx_name( + "vllm_internal_worker_extension/update_weights_from_sparse_payload_files" + ) + def update_weights_from_sparse_payload_files( + self, *payload_paths: str ) -> dict[str, Any]: - locations, values, metadata = payload - - sparse_started = time.perf_counter() - self._apply_sparse_weight_deltas((locations, values), metadata) - sparse_apply_s = time.perf_counter() - sparse_started - - return { - "ok": True, - "receiver_sparse_apply_s": sparse_apply_s, - } + return self._load_and_apply_sparse_payloads(payload_paths) def synchronize_device(self) -> None: - """Synchronize this vLLM worker's CUDA device after deferred refit applies.""" if torch.cuda.is_available(): torch.cuda.synchronize(self._cuda_device_index) def finish_sparse_delta_refit(self) -> dict[str, Any]: """Synchronize and compare bounded producer samples with applied weights.""" self.synchronize_device() - verification = self._direct_sparse_delta_verification or [] - candidates = self._direct_sparse_delta_verification_candidates - self._direct_sparse_delta_verification = [] - self._direct_sparse_delta_verification_candidates = 0 + verification, self._verification = self._verification, [] + candidates, self._verification_candidates = self._verification_candidates, 0 if not verification: return { "ok": True, @@ -730,30 +389,28 @@ def finish_sparse_delta_refit(self) -> dict[str, Any]: } with torch.no_grad(): - actual_delta = torch.cat( + actual = torch.cat( [ target.data.view(-1).index_select(0, locations).float() - before for target, locations, before, _ in verification ] ) - expected_delta = torch.cat([expected for _, _, _, expected in verification]) - difference = (actual_delta - expected_delta).abs() - exact_mismatches = actual_delta.ne(expected_delta) - mismatches = ~torch.isclose( - actual_delta, expected_delta, rtol=1e-6, atol=1e-8 - ) + expected = torch.cat([item[3] for item in verification]) + difference = (actual - expected).abs() stats = torch.stack( ( difference.sum(), difference.max(), - exact_mismatches.sum().float(), - mismatches.sum().float(), + actual.ne(expected).sum().float(), + (~torch.isclose(actual, expected, rtol=1e-6, atol=1e-8)) + .sum() + .float(), ) ).cpu() return { "ok": True, "verification_candidates": candidates, - "verification_samples": actual_delta.numel(), + "verification_samples": actual.numel(), "verification_exact_mismatches": int(stats[2]), "verification_mismatches": int(stats[3]), "verification_abs_sum": float(stats[0]), diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py index 5d3050ca8e4..68d35817bea 100644 --- a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py +++ b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py @@ -38,7 +38,7 @@ G_VLLM_REFIT_S3_MANIFEST_PATH, decode_sparse_payload, download_s3_refit_payload, - merge_vllm_refit_receiver_timing, + merge_vllm_refit_metrics, refit_env_int, vllm_refit_api_key, ) @@ -79,19 +79,11 @@ def __init__(self, worker: Any) -> None: self._zmq_refit_server: tuple[ZmqSparseRefitServer, str] | None = None self._refit_async_loop: asyncio.AbstractEventLoop | None = None - @property - def cfg(self) -> Any: - return self._worker.cfg - - @property - def llm(self) -> Any: - return self._worker.llm - def set_worker_hostnames(self, hostnames: list[str]) -> None: self._refit_workers_share_node = len(set(hostnames)) == 1 def start_sync_server(self) -> None: - llm = self.llm + llm = self._worker.llm if llm is None: raise RuntimeError("vLLM is not initialized on this worker.") self.set_worker_hostnames(llm.collective_rpc("report_node_hostname", args=())) @@ -161,7 +153,7 @@ def _collect_refit_apply_results( ) -> dict[str, Any]: results = [future.result() for future in futures] timing: dict[str, float] = {} - merge_vllm_refit_receiver_timing(timing, results, maximum=False) + merge_vllm_refit_metrics(timing, results, maximum=False) return { "ok": True, "payloads": sum(int(result.get("payloads", 0)) for result in results), @@ -171,37 +163,22 @@ def _collect_refit_apply_results( @staticmethod def _refit_collective_response(worker_results: Any) -> dict[str, Any]: results = cast(list[dict[str, Any]], worker_results) - response = { + return { "ok": True, - **merge_vllm_refit_receiver_timing({}, results, maximum=True), + **merge_vllm_refit_metrics( + {}, results, maximum=True, candidate_maximum=True + ), } - if any("verification_candidates" in result for result in results): - response["verification_candidates"] = max( - (int(result["verification_candidates"]) for result in results), - default=0, - ) - for key in ( - "verification_samples", - "verification_exact_mismatches", - "verification_mismatches", - "verification_abs_sum", - ): - response[key] = sum(result[key] for result in results) - response["verification_max_abs"] = max( - (float(result["verification_max_abs"]) for result in results), - default=0.0, - ) - return response def _refit_collective_rpc( self, method: str, args: tuple[Any, ...], ) -> Any: - llm = self.llm + llm = self._worker.llm if llm is None: raise RuntimeError("vLLM is not initialized on this worker.") - if not self.cfg["vllm_cfg"]["async_engine"]: + if not self._worker.cfg["vllm_cfg"]["async_engine"]: return llm.collective_rpc(method, args=args) if self._refit_async_loop is None: raise RuntimeError("The async vLLM refit server loop is not initialized.") @@ -215,21 +192,15 @@ def update_weights_from_serialized_sparse_payloads( serialized_payloads: tuple[bytes, ...], ) -> dict[str, Any]: """Apply a FIFO batch of sparse deltas through one collective RPC.""" - if self.llm is None: - raise RuntimeError("vLLM is not initialized on this worker.") if not self._refit_workers_share_node: - results = [ - self._refit_collective_response( - self._refit_collective_rpc( - "update_weights_from_serialized_sparse_payload", - (payload,), - ) + response = self._refit_collective_response( + self._refit_collective_rpc( + "update_weights_from_serialized_sparse_payload", + serialized_payloads, ) - for payload in serialized_payloads - ] - timing: dict[str, float] = {} - merge_vllm_refit_receiver_timing(timing, results, maximum=False) - return {"ok": True, "payloads": len(serialized_payloads), **timing} + ) + response["payloads"] = len(serialized_payloads) + return response with tempfile.TemporaryDirectory( prefix="nemo_rl_refit_", dir=self._refit_batch_staging_dir @@ -271,7 +242,6 @@ def _flush_queued_sparse_payloads(self) -> dict[str, Any]: submitted.add_done_callback(self._notify_refit_apply_waiters) response = self._collect_refit_apply_results(futures) if futures: - assert self.llm is not None response.update( self._refit_collective_response( self._refit_collective_rpc("finish_sparse_delta_refit", ()) @@ -347,15 +317,14 @@ async def _apply_zmq_payload(self, raw_request: Any) -> dict[str, Any]: return result def setup_api_server(self, app: Any) -> None: - token = vllm_refit_api_key( - self.cfg["vllm_cfg"].get("http_refit_api_key_env_var") - ) + cfg = self._worker.cfg + token = vllm_refit_api_key(cfg["vllm_cfg"].get("http_refit_api_key_env_var")) async def respond( raw_request: Request, action: Literal["s3", "flush", "zmq"], ) -> JSONResponse: - if self.cfg["vllm_cfg"]["async_engine"]: + if cfg["vllm_cfg"]["async_engine"]: self._refit_async_loop = asyncio.get_running_loop() if ( token is not None @@ -398,16 +367,15 @@ def report_refit_server_base_url(self) -> str | None: def start_zmq_sparse_refit_relay(self, refit_urls: list[str]) -> str: if self._zmq_refit_server is not None: return self._zmq_refit_server[1] - port = self.cfg["vllm_cfg"].get( - "zmq_refit_server_port" - ) or _get_free_port_local( - self.cfg.get("port_range_low", DEFAULT_GENERATION_PORT_RANGE_LOW), - self.cfg.get("port_range_high", DEFAULT_GENERATION_PORT_RANGE_HIGH), + cfg = self._worker.cfg + port = cfg["vllm_cfg"].get("zmq_refit_server_port") or _get_free_port_local( + cfg.get("port_range_low", DEFAULT_GENERATION_PORT_RANGE_LOW), + cfg.get("port_range_high", DEFAULT_GENERATION_PORT_RANGE_HIGH), ) server = ZmqSparseRefitServer( refit_urls, bind_address=f"tcp://0.0.0.0:{port}", - api_key_env_var=self.cfg["vllm_cfg"].get("http_refit_api_key_env_var"), + api_key_env_var=cfg["vllm_cfg"].get("http_refit_api_key_env_var"), timeout_s=float(os.getenv("NRL_REFIT_ZMQ_TIMEOUT_S") or 600.0), ) server.start() @@ -424,11 +392,10 @@ def stop_zmq_sparse_refit_relay(self) -> None: def _setup_vllm_refit_server(self) -> None: app = FastAPI() self.setup_api_server(app) - port = self.cfg["vllm_cfg"].get( - "http_refit_server_port" - ) or _get_free_port_local( - self.cfg.get("port_range_low", DEFAULT_GENERATION_PORT_RANGE_LOW), - self.cfg.get("port_range_high", DEFAULT_GENERATION_PORT_RANGE_HIGH), + cfg = self._worker.cfg + port = cfg["vllm_cfg"].get("http_refit_server_port") or _get_free_port_local( + cfg.get("port_range_low", DEFAULT_GENERATION_PORT_RANGE_LOW), + cfg.get("port_range_high", DEFAULT_GENERATION_PORT_RANGE_HIGH), ) server = uvicorn.Server( uvicorn.Config( diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 32f284fe740..28b49cb0eac 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -366,7 +366,7 @@ def __init__( self.sampling_params = runtime_config.sampling_params generation_config = self.cfg.get("generation") delta_config = None - if generation_config is not None and generation_config["backend"] == "vllm": + if generation_config and generation_config.get("refit_transport") is not None: delta_config = cast(VllmConfig, generation_config).get("delta_compression") self._remote_sparse_refit = None if delta_config: diff --git a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py index 895f143999c..b3e36deeeac 100644 --- a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py +++ b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py @@ -61,11 +61,7 @@ def stream( streamer = { "s3": stream_sparse_delta_payloads_via_s3_manifest, "zmq": stream_sparse_delta_payloads_via_zmq, - }.get(transport) - if streamer is None: - raise ValueError( - f"Unsupported remote sparse refit transport {transport!r}." - ) + }[transport] result = streamer( self._worker._iter_params_with_optional_kv_scales(), delta_tracker=self._tracker, diff --git a/nemo_rl/utils/weight_transfer_remote_sparse.py b/nemo_rl/utils/weight_transfer_remote_sparse.py index e2be67a8839..b4f2c28faa9 100644 --- a/nemo_rl/utils/weight_transfer_remote_sparse.py +++ b/nemo_rl/utils/weight_transfer_remote_sparse.py @@ -80,7 +80,6 @@ def _s3_client(region: str) -> Any: class _S3ObjectStore: def __init__(self, *, bucket: str, region: str) -> None: - # Keep the AWS runtime unloaded for ZeroMQ-only jobs. from awscrt.s3 import S3RequestType self.bucket = bucket @@ -95,7 +94,6 @@ def put(self, key: str, body: bytes) -> None: ).finished_future.result() def get(self, key: str) -> bytes: - # Keep the AWS runtime unloaded for ZeroMQ-only jobs. from awscrt.http import HttpHeaders body = bytearray() @@ -132,7 +130,6 @@ def delete(self, key: str) -> None: ).finished_future.result() def _request(self, method: str, key: str, body: bytes | None = None) -> Any: - # Keep the AWS runtime unloaded for ZeroMQ-only jobs. from awscrt.http import HttpHeaders, HttpRequest headers = HttpHeaders( @@ -303,7 +300,6 @@ def stream_sparse_delta_payloads( f"NRL_REFIT_{prefix}_ENCODE_WORKERS", default=max(2, min(8, os.cpu_count() or 8)), ) - pipeline_workers = max(encode_workers, transfer_workers) encode_executor = _executor(f"refit-{transport}-encode", encode_workers) transfer_executor = _executor(f"refit-{transport}-transfer", transfer_workers) export_chunk_size = sparse_export_chunk_size(delta_tracker, transport) @@ -382,7 +378,7 @@ def collect_transfers(*, block: bool) -> None: for key, value in result.items(): if key.endswith("_s"): timing[key] = timing.get(key, 0.0) + float(value) - merge_vllm_refit_receiver_timing( + merge_vllm_refit_metrics( receiver_timing, [result["receiver"]], maximum=False ) @@ -445,7 +441,6 @@ def drain_encodes() -> None: "payloads": counts["payloads"], "chunks": chunk_count, "wire_mb": counts["wire_bytes"] / 1e6, - "pipeline_workers": pipeline_workers, "encode_workers": encode_workers, "export_chunk_mb": export_chunk_size / 1e6, "shard_rank": shard_rank, @@ -524,7 +519,7 @@ def send_payload(body: bytes, payload_index: int) -> dict[str, Any]: return { "s3_put_s": s3_put_s, "manifest_post_s": manifest_post_s, - "receiver": merge_vllm_refit_receiver_timing({}, responses, maximum=True), + "receiver": merge_vllm_refit_metrics({}, responses, maximum=True), } return stream_sparse_delta_payloads( @@ -603,23 +598,29 @@ def download_s3_refit_payload( return decode_sparse_payload(body, checksum) -def merge_vllm_refit_receiver_timing( +def merge_vllm_refit_metrics( result: dict[str, Any], - timings: Iterable[Mapping[str, Any]], + metrics: Iterable[Mapping[str, Any]], *, maximum: bool, + candidate_maximum: bool | None = None, ) -> dict[str, Any]: - for timing in timings: - for key, value in timing.items(): + for metric in metrics: + for key, value in metric.items(): if key.startswith("receiver_") and key.endswith("_s"): - number = float(value) - if key in result: - number = ( - max(float(result[key]), number) - if maximum - else float(result[key]) + number - ) - result[key] = number + number, use_maximum = float(value), maximum + elif candidate_maximum is not None and key.startswith("verification_"): + number = value + use_maximum = key == "verification_max_abs" or ( + key == "verification_candidates" and candidate_maximum + ) + else: + continue + if key in result: + number = ( + max(result[key], number) if use_maximum else result[key] + number + ) + result[key] = number return result diff --git a/nemo_rl/utils/weight_transfer_sparse_codec.py b/nemo_rl/utils/weight_transfer_sparse_codec.py index 2786f1ef527..aaaf7fa2685 100644 --- a/nemo_rl/utils/weight_transfer_sparse_codec.py +++ b/nemo_rl/utils/weight_transfer_sparse_codec.py @@ -107,10 +107,12 @@ def _encode_explicit_locations( indices = locations.detach().cpu().numpy().astype(np.int64, copy=False) deltas = np.diff(indices, prepend=-1) - 1 max_delta = int(deltas.max()) - dtype = next( - dtype - for dtype in (np.uint16, np.uint32, np.uint64) - if max_delta <= np.iinfo(dtype).max + dtype = ( + np.uint16 + if max_delta <= 0xFFFF + else np.uint32 + if max_delta <= 0xFFFFFFFF + else np.uint64 ) raw = deltas.astype(dtype, copy=False).tobytes() return torch.from_numpy(np.frombuffer(raw, dtype=np.uint8).copy()) @@ -193,9 +195,6 @@ def _add_verification_samples( ) -> None: total = sum(int(locations.numel()) for locations, _ in sources) count = min(self.verification_samples, total) - if not count: - return - sample_ranks = [ ((2 * index + 1) * total) // (2 * count) for index in range(count) ] diff --git a/nemo_rl/utils/weight_transfer_zmq.py b/nemo_rl/utils/weight_transfer_zmq.py index f43999a9f83..c6b4f151330 100644 --- a/nemo_rl/utils/weight_transfer_zmq.py +++ b/nemo_rl/utils/weight_transfer_zmq.py @@ -26,7 +26,7 @@ from nemo_rl.utils.weight_transfer_remote_sparse import ( SparseDeltaStreamResult, - merge_vllm_refit_receiver_timing, + merge_vllm_refit_metrics, post_vllm_refit_endpoints, refit_env_int, sparse_payload_checksum, @@ -231,7 +231,7 @@ def _fanout( headers=headers, executor=http_executor, ) - merged = merge_vllm_refit_receiver_timing({}, results, maximum=True) + merged = merge_vllm_refit_metrics({}, results, maximum=True) merged["receiver_relay_fanout_s"] = time.perf_counter() - started return merged @@ -246,7 +246,7 @@ def _send_reply( socket.send_multipart( [identity, kind, _json_bytes(reply)], flags=zmq.NOBLOCK ) - except (zmq.Again, zmq.ZMQError): + except zmq.ZMQError: pass def _parse_data_message( diff --git a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py index eab7fab6ad9..c16969ae47b 100644 --- a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py +++ b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py @@ -16,13 +16,17 @@ import time import uuid +from collections import defaultdict from contextlib import nullcontext, suppress from typing import Any import ray from nemo_rl.utils.timer import Timer -from nemo_rl.utils.weight_transfer_remote_sparse import flush_vllm_refit_urls +from nemo_rl.utils.weight_transfer_remote_sparse import ( + flush_vllm_refit_urls, + merge_vllm_refit_metrics, +) from nemo_rl.weight_sync.interfaces import WeightSynchronizer @@ -66,13 +70,13 @@ def __init__( self._policy = policy self._generation = generation self._transport = transport - self._refit_urls: list[str] = [] - self._targets: list[str] = [] self._api_key_env_var = api_key_env_var self._request_timeout_s = request_timeout_s + self._refit_urls: list[str] = [] + self._targets: list[str] = [] + self._baseline_init_refs: list[Any] = [] + self._baseline_commit_refs: list[Any] = [] self._stale = True - self._baseline_init_refs: list[Any] | None = None - self._baseline_commit_refs: list[Any] | None = None def sync_weights( self, @@ -80,94 +84,75 @@ def sync_weights( timer: Timer | None = None, kv_scales: dict[str, float] | None = None, ) -> dict[str, float]: - timer_context = ( + context = ( timer.time("prepare_for_generation/transfer_and_update_weights") - if timer is not None + if timer else nullcontext() ) - with timer_context: - if self._baseline_commit_refs is not None: + with context: + if self._baseline_commit_refs: ray.get(self._baseline_commit_refs) - self._baseline_commit_refs = None + self._baseline_commit_refs.clear() if not self._generation.invalidate_kv_cache(): raise RuntimeError( f"vLLM KV cache invalidation failed before {self._transport} " "weight update." ) - - if self._baseline_init_refs is not None: + if self._baseline_init_refs: ray.get(self._baseline_init_refs) - self._baseline_init_refs = None + self._baseline_init_refs.clear() + succeeded = False try: transfer_id = uuid.uuid4().hex - refs = self._run_policy_workers( - "stream_remote_sparse_weights", - transport=self._transport, - targets=self._targets, - transfer_id=transfer_id, - api_key_env_var=self._api_key_env_var, - timeout_s=self._request_timeout_s, + results = ray.get( + self._run_policy_workers( + "stream_remote_sparse_weights", + transport=self._transport, + targets=self._targets, + transfer_id=transfer_id, + api_key_env_var=self._api_key_env_var, + timeout_s=self._request_timeout_s, + ) ) - results = ray.get(refs) payloads = sum(result["payloads"] for result in results) - changed_elements = sum(result["changed_elements"] for result in results) - total_elements = sum(result["total_elements"] for result in results) - changed_pct = 100.0 * changed_elements / max(total_elements, 1) + changed = sum(result["changed_elements"] for result in results) + total = sum(result["total_elements"] for result in results) + changed_pct = 100.0 * changed / max(total, 1) print( f"REFIT_{self._transport.upper()}_DELTA_CHANGE " - f"changed_elements={changed_elements} " - f"total_elements={total_elements} " + f"changed_elements={changed} total_elements={total} " f"changed_pct={changed_pct:.8g}", flush=True, ) - candidates = 0 - samples = 0 - exact_mismatches = 0 - mismatches = 0 - abs_sum = 0.0 - max_abs = 0.0 + + verification: defaultdict[str, float] = defaultdict(float) commit_s = 0.0 if payloads: started = time.perf_counter() - receiver_results = flush_vllm_refit_urls( - self._refit_urls, - api_key_env_var=self._api_key_env_var, - timeout_s=self._request_timeout_s, - ) - candidates = sum( - int(result.get("verification_candidates", 0)) - for result in receiver_results - ) - samples = sum( - int(result.get("verification_samples", 0)) - for result in receiver_results - ) - exact_mismatches = sum( - int(result.get("verification_exact_mismatches", 0)) - for result in receiver_results - ) - mismatches = sum( - int(result.get("verification_mismatches", 0)) - for result in receiver_results - ) - abs_sum = sum( - float(result.get("verification_abs_sum", 0.0)) - for result in receiver_results - ) - max_abs = max( - ( - float(result.get("verification_max_abs", 0.0)) - for result in receiver_results - ), - default=0.0, + verification.update( + merge_vllm_refit_metrics( + {}, + flush_vllm_refit_urls( + self._refit_urls, + api_key_env_var=self._api_key_env_var, + timeout_s=self._request_timeout_s, + ), + maximum=True, + candidate_maximum=False, + ) ) + candidates = int(verification["verification_candidates"]) + samples = int(verification["verification_samples"]) + exact = int(verification["verification_exact_mismatches"]) + mismatches = int(verification["verification_mismatches"]) + abs_sum = float(verification["verification_abs_sum"]) + max_abs = float(verification["verification_max_abs"]) if candidates or samples: print( f"REFIT_{self._transport.upper()}_DELTA_VERIFY " f"candidates={candidates} samples={samples} " - f"exact_mismatches={exact_mismatches} " - f"mismatches={mismatches} " + f"exact_mismatches={exact} mismatches={mismatches} " f"mean_abs={abs_sum / max(samples, 1):.8g} " f"max_abs={max_abs:.8g}", flush=True, @@ -195,25 +180,36 @@ def sync_weights( ) self._baseline_commit_refs = ( self._policy.worker_group.run_all_workers_single_data( - "finish_remote_sparse_delta_sync", - succeeded=succeeded, + "finish_remote_sparse_delta_sync", succeeded=succeeded ) ) + self._stale = False - return { - "delta/changed_elements": float(changed_elements), - "delta/total_elements": float(total_elements), + samples = int(verification["verification_samples"]) + mismatches = int(verification["verification_mismatches"]) + metrics = { + "delta/changed_elements": float(changed), + "delta/total_elements": float(total), "delta/changed_pct": changed_pct, - "delta_verify/candidates": float(candidates), - "delta_verify/samples": float(samples), - "delta_verify/exact_mismatches": float(exact_mismatches), - "delta_verify/mismatches": float(mismatches), "delta_verify/mismatch_pct": 100.0 * mismatches / max(samples, 1), - "delta_verify/mean_abs": abs_sum / max(samples, 1), - "delta_verify/max_abs": max_abs, + "delta_verify/mean_abs": float(verification["verification_abs_sum"]) + / max(samples, 1), "transfer/payloads": float(payloads), "transfer/global_commit_s": commit_s, } + metrics.update( + { + f"delta_verify/{key}": float(verification[f"verification_{key}"]) + for key in ( + "candidates", + "samples", + "exact_mismatches", + "mismatches", + "max_abs", + ) + } + ) + return metrics @property def is_stale(self) -> bool: @@ -223,20 +219,18 @@ def mark_stale(self) -> None: self._stale = True def _run_policy_workers(self, method_name: str, **kwargs: Any) -> list[Any]: - worker_group = self._policy.worker_group - worker_count = len(worker_group.workers) - return worker_group.run_all_workers_multiple_data( + workers = self._policy.worker_group + count = len(workers.workers) + return workers.run_all_workers_multiple_data( method_name, - common_kwargs={**kwargs, "shard_count": worker_count}, - shard_rank=list(range(worker_count)), + common_kwargs={**kwargs, "shard_count": count}, + shard_rank=list(range(count)), ) def _run_generation_workers(self, method_name: str, **kwargs: Any) -> list[Any]: - worker_group = self._generation.worker_group - if not worker_group or not worker_group.workers: - raise RuntimeError("vLLM worker group is not initialized.") + workers = self._generation.worker_group return ray.get( - worker_group.run_all_workers_single_data( + workers.run_all_workers_single_data( method_name, run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], **kwargs, @@ -245,8 +239,7 @@ def _run_generation_workers(self, method_name: str, **kwargs: Any) -> list[Any]: def init_communicator(self) -> None: self._baseline_init_refs = self._run_policy_workers( - "init_remote_sparse_delta_baseline", - transport=self._transport, + "init_remote_sparse_delta_baseline", transport=self._transport ) self._refit_urls = [ url @@ -269,14 +262,12 @@ def init_communicator(self) -> None: self._stale = False def shutdown(self) -> None: - for ref in (self._baseline_init_refs or []) + ( - self._baseline_commit_refs or [] - ): + for ref in self._baseline_init_refs + self._baseline_commit_refs: ray.cancel(ref, force=False) if self._transport == "zmq": self._run_generation_workers("stop_zmq_sparse_refit_relay") - self._baseline_init_refs = None - self._baseline_commit_refs = None - self._refit_urls = [] - self._targets = [] + self._baseline_init_refs.clear() + self._baseline_commit_refs.clear() + self._refit_urls.clear() + self._targets.clear() self._stale = True diff --git a/tests/unit/models/generation/test_vllm_sparse_delta.py b/tests/unit/models/generation/test_vllm_sparse_delta.py index f23202ea822..49b2c2bcbcb 100644 --- a/tests/unit/models/generation/test_vllm_sparse_delta.py +++ b/tests/unit/models/generation/test_vllm_sparse_delta.py @@ -12,277 +12,299 @@ # See the License for the specific language governing permissions and # limitations under the License. +import math +import sys from types import MethodType, SimpleNamespace from typing import Any import pytest import torch -from nemo_rl.models.generation.vllm.vllm_sparse_delta import VllmSparseDeltaApplier +from nemo_rl.models.generation.vllm.vllm_sparse_delta import ( + VllmSparseDeltaApplier, + _SparseLoadTracer, +) from nemo_rl.utils.weight_transfer_sparse_codec import encode_sparse_infos -def _attach_tensor_attrs(tensor: torch.Tensor, **attrs: object) -> torch.Tensor: - for name, value in attrs.items(): - setattr(tensor, name, value) - return tensor +class _NativeLoaderModel: + def __init__(self, **targets: torch.Tensor) -> None: + self.targets = targets + def parameters(self): + return iter(self.targets.values()) -def _make_sparse_delta_extension( - parameter_name: str, - target: torch.Tensor, - module: object, -) -> Any: - model_runner = SimpleNamespace( - model=SimpleNamespace(get_submodule=lambda _name: module) - ) - ext = VllmSparseDeltaApplier( - model_runner, + def buffers(self): + return iter(()) + + def load_weights(self, weights): + for name, source in weights: + target = self.targets + if name == "weight": + target["identity"].copy_(source) + elif name.endswith("self_attn.k_proj.weight"): + target["qkv"][4:6].copy_(source[2:4]) + elif name.endswith("mlp.gate_proj.weight"): + target["merged"][:4].copy_(source[4:8]) + elif name.endswith("mlp.up_proj.weight"): + target["merged"][4:8].copy_(source[4:8]) + elif name.endswith("experts.3.gate_proj.weight"): + target["w13"][1, :4].copy_(source[4:8]) + elif name.endswith("experts.3.down_proj.weight"): + target["w2"][1].copy_(source[:, 4:8]) + elif name.endswith("mixer.in_proj.weight"): + target["mamba"][:2].copy_(source[2:4]) + target["mamba"][2:].copy_(source[6:10]) + elif name.endswith(("mixer.A", "mixer.A_log")): + target["a"].copy_(source) + target["a"].copy_(-torch.exp(target["a"])) + elif name == "transformed": + target["identity"].copy_(source + 1) + return set() + + +def _applier(model: Any) -> VllmSparseDeltaApplier: + return VllmSparseDeltaApplier( + SimpleNamespace( + model=model, + vllm_config=SimpleNamespace(), + ), torch.device("cpu"), - rank=1, ) - ext._direct_sparse_delta_targets = {parameter_name: target} - return ext -def _assert_sparse_plan( - ext: Any, - plan: Any, - source_locations: list[int], - expected_locations: list[int], - expected_values: list[float], -) -> None: - assert plan is not None - values = torch.arange(len(source_locations), dtype=torch.float32) - locations, values = ext._local_sparse_delta_update_inputs( - torch.tensor(source_locations), values, plan +def _stub_fp8(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setitem( + sys.modules, + "nemo_rl.models.generation.vllm.quantization.fp8", + SimpleNamespace(is_fp8_model=lambda _config: False), ) - assert locations.tolist() == expected_locations - assert values.tolist() == expected_values @pytest.mark.vllm def test_serialized_sparse_payload_batch_preserves_order(tmp_path) -> None: - ext = VllmSparseDeltaApplier(SimpleNamespace(), torch.device("cpu")) + applier = VllmSparseDeltaApplier(SimpleNamespace(), torch.device("cpu")) payloads = [ - (torch.tensor([index]), torch.tensor([float(index)]), {"index": index}) + (torch.tensor([index]), torch.tensor([float(index)]), [{"index": index}]) for index in range(3) ] paths = [tmp_path / f"{index}.pt" for index in range(3)] for path, payload in zip(paths, payloads, strict=True): torch.save(payload, path) applied: list[Any] = [] + compiled: list[Any] = [] + applier._compile_plans = compiled.append + applier._apply_sparse_weight_deltas = lambda tensors, metadata: applied.append( + (*tensors, metadata) + ) - def apply(payload: Any) -> dict[str, Any]: - applied.append(payload) - return { - "ok": True, - "receiver_sparse_apply_s": 2.0, - } - - ext._apply_sparse_request = apply - result = ext.update_weights_from_sparse_payload_files( + result = applier.update_weights_from_sparse_payload_files( *(str(path) for path in paths) ) - - assert [item[2]["index"] for item in applied] == [0, 1, 2] - assert all( - torch.equal(item[1], payload[1]) - for item, payload in zip(applied, payloads, strict=True) + applier.update_weights_from_serialized_sparse_payload( + *(path.read_bytes() for path in paths) ) + + assert [[item["index"] for item in batch] for batch in compiled] == [[0, 1, 2]] * 2 + assert [item[2][0]["index"] for item in applied] == [0, 1, 2] * 2 assert result["receiver_deserialize_s"] >= 0.0 - assert result["receiver_sparse_apply_s"] == 6.0 + assert result["receiver_plan_s"] >= 0.0 + assert result["receiver_sparse_apply_s"] >= 0.0 @pytest.mark.vllm -def test_direct_sparse_delta_placement() -> None: - qkv_name = "model.layers.0.self_attn.qkv_proj.weight" - qkv_target = _attach_tensor_attrs(torch.zeros(8, 2), output_dim=0) - ext = _make_sparse_delta_extension( - qkv_name, - qkv_target, - SimpleNamespace( - tp_rank=1, - num_kv_head_replicas=2, - _get_shard_offset_mapping=lambda shard: {"q": 0, "k": 4, "v": 6}[shard], - _get_shard_size_mapping=lambda shard: {"q": 4, "k": 2, "v": 2}[shard], +def test_native_loaders_compile_sparse_placement(monkeypatch) -> None: + _stub_fp8(monkeypatch) + targets = { + "identity": torch.zeros(4), + "qkv": torch.zeros(8, 2), + "merged": torch.zeros(8, 2), + "w13": torch.zeros(2, 4, 2), + "w2": torch.zeros(2, 2, 4), + "mamba": torch.zeros(6, 2), + "a": torch.tensor([-2.0, -4.0]), + } + infos = [ + ("weight", (4,), [1, 2], [1.0, 2.0]), + ("model.layers.0.self_attn.k_proj.weight", (4, 2), [0, 4, 5, 7], [1] * 4), + ("model.layers.0.mlp.gate_proj.weight", (8, 2), [0, 8, 9, 15], [2] * 4), + ("model.layers.0.mlp.up_proj.weight", (8, 2), [8, 15], [3] * 2), + ("model.layers.0.mlp.experts.3.gate_proj.weight", (8, 2), [8, 15], [4] * 2), + ( + "model.layers.0.mlp.experts.3.down_proj.weight", + (2, 8), + [3, 4, 7, 12, 15], + [5] * 5, + ), + ("model.layers.0.mixer.in_proj.weight", (10, 2), [0, 4, 7, 12, 19], [6] * 5), + ( + "backbone.layers.0.mixer.A_log", + (2,), + [0, 1], + [math.log(1.5), math.log(0.5)], ), + ("model.layers.0.mlp.experts.7.gate_proj.weight", (8, 2), [8], [9]), + ] + payload = encode_sparse_infos( + [ + ( + name, + torch.empty(shape), + torch.tensor(locations), + torch.tensor(values, dtype=torch.float32), + ) + for name, shape, locations, values in infos + ], + empty_dtype=torch.float32, ) - qkv_source = "model.layers.0.self_attn.k_proj.weight" - plan = ext._direct_sparse_delta_qkv_plan( - {"name": qkv_source, "shape": (2, 2)}, qkv_source, {qkv_name: qkv_target} + payload[2][-1].update(verification_locations=[8], verification_deltas=[9.0]) + + applier = _applier(_NativeLoaderModel(**targets)) + applier._apply_sparse_weight_deltas(payload[:2], payload[2]) + + assert torch.equal(targets["identity"], torch.tensor([0.0, 1.0, 2.0, 0.0])) + assert targets["qkv"].view(-1)[[8, 9, 11]].tolist() == [1.0, 1.0, 1.0] + assert targets["merged"].view(-1)[[0, 1, 7, 8, 15]].tolist() == [2, 2, 2, 3, 3] + assert targets["w13"].view(-1)[[8, 15]].tolist() == [4, 4] + assert targets["w2"].view(-1)[[8, 11, 12, 15]].tolist() == [5, 5, 5, 5] + assert targets["mamba"].view(-1)[[0, 3, 4, 11]].tolist() == [6, 6, 6, 6] + assert torch.allclose(targets["a"], torch.tensor([-3.0, -2.0])) + assert (applier._verification_candidates, applier._verification) == (1, []) + + +@pytest.mark.vllm +def test_vllm_native_loader_geometry() -> None: + from vllm.model_executor.layers.fused_moe.layer import FusedMoE + from vllm.model_executor.layers.linear import ( + MergedColumnParallelLinear, + QKVParallelLinear, ) - _assert_sparse_plan(ext, plan, [0, 1, 2, 3], [8, 9, 10, 11], [0.0, 1.0, 2.0, 3.0]) - - merged_name = "model.layers.0.mlp.gate_up_proj.weight" - merged_target = _attach_tensor_attrs(torch.zeros(8, 2), output_dim=0) - ext = _make_sparse_delta_extension( - merged_name, - merged_target, - SimpleNamespace(tp_rank=1, tp_size=2, output_sizes=(8, 8)), + from vllm.model_executor.layers.mamba.mamba_mixer2 import ( + mamba_v2_sharded_weight_loader, ) - for projection, expected_locations in ( - ("gate", [0, 1, 6, 7]), - ("up", [8, 9, 14, 15]), - ): - source_name = f"model.layers.0.mlp.{projection}_proj.weight" - plan = ext._direct_sparse_delta_target_plan( - {"name": source_name, "shape": (8, 2)}, - {merged_name: merged_target}, - ) - _assert_sparse_plan( - ext, - plan, - [6, 7, 8, 9, 14, 15], - expected_locations, - [2.0, 3.0, 4.0, 5.0], - ) - expert_name = "model.layers.0.mlp.experts.w13_weight" - expert_target = torch.zeros(2, 4, 2) - expert_module = SimpleNamespace( - tp_rank=1, - moe_config=SimpleNamespace(is_act_and_mul=False), - _map_global_expert_id_to_local_expert_id=lambda expert: ( - 1 if expert == 3 else -1 + def trace(target, source_shape, load): + source = torch.empty(source_shape, device="meta") + tracer = _SparseLoadTracer([target], {"source": source}) + with tracer: + load(source) + return tracer.copies["source"] + + qkv = torch.nn.Parameter(torch.zeros(8, 2)) + qkv.output_dim = 0 + qkv_layer = SimpleNamespace( + num_heads=4, + num_kv_heads=2, + num_kv_head_replicas=2, + head_size=1, + v_head_size=1, + tp_rank=3, + ) + qkv_layer.validate_shard_id = MethodType( + QKVParallelLinear.validate_shard_id, qkv_layer + ) + qkv_copies = trace( + qkv, + (4, 2), + lambda source: QKVParallelLinear.weight_loader(qkv_layer, qkv, source, "k"), + ) + + merged = torch.nn.Parameter(torch.zeros(8, 2)) + merged.output_dim = 0 + merged_layer = SimpleNamespace(output_sizes=[8, 8], tp_size=2, tp_rank=1) + merged_layer.validate_shard_id = MethodType( + MergedColumnParallelLinear.validate_shard_id, merged_layer + ) + merged_copies = trace( + merged, + (8, 2), + lambda source: MergedColumnParallelLinear.weight_loader( + merged_layer, merged, source, 1 ), ) - ext = _make_sparse_delta_extension( - expert_name, - expert_target, - expert_module, + + expert = torch.nn.Parameter(torch.zeros(2, 8, 2)) + moe = SimpleNamespace( + moe_config=SimpleNamespace(is_act_and_mul=True), + _get_hidden_dim=FusedMoE._get_hidden_dim, + _narrow_expert_data_for_padding=FusedMoE._narrow_expert_data_for_padding, + ) + expert_copies = trace( + expert, + (8, 2), + lambda source: FusedMoE._load_w13(moe, expert.data[1], 0, "w3", source, 1), ) - for projection in ("gate_proj", "up_proj"): - expert_source = f"model.layers.0.mlp.experts.3.{projection}.weight" - plan = ext._direct_sparse_delta_expert_plan( - {"name": expert_source, "shape": (8, 2)}, - expert_source, - {expert_name: expert_target}, - ) - _assert_sparse_plan( - ext, - plan, - [6, 7, 8, 9, 14, 15], - [8, 9, 14, 15], - [2.0, 3.0, 4.0, 5.0], - ) - w2_target = torch.zeros(2, 2, 4) - ext = _make_sparse_delta_extension(expert_name, w2_target, expert_module) - expert_source = "model.layers.0.mlp.experts.3.down_proj.weight" - plan = ext._direct_sparse_delta_expert_plan( - {"name": expert_source, "shape": (2, 8)}, - expert_source, - {"model.layers.0.mlp.experts.w2_weight": w2_target}, + mamba = torch.nn.Parameter(torch.zeros(10, 1, 2)) + mamba_loader = mamba_v2_sharded_weight_loader( + [(8, 0, False), (4, 2, True), (4, 2, True), (4, 0, False)], 2, 1 ) - _assert_sparse_plan(ext, plan, [3, 4, 7, 11, 15], [8, 11, 15], [1.0, 2.0, 4.0]) + mamba_copies = trace(mamba, (16, 1, 2), lambda source: mamba_loader(mamba, source)) - mamba_name = "model.layers.0.mixer.in_proj.weight" - for target_shape, groups, source_locations, expected_locations, values in ( - ((16, 1, 2), 6, [0, 8, 25, 38, 40, 55], [0, 9, 22, 31], [1, 2, 4, 5]), + cases = ( + (qkv_copies, [0, 4, 7], [8, 11]), + (merged_copies, [0, 8, 15], [8, 15]), + (expert_copies, [0, 8, 15], [24, 31]), ( - (14, 2), - 4, - [0, 8, 24, 36, 44, 52, 55], - [0, 8, 16, 20, 24, 27], - [1, 2, 3, 4, 5, 6], + mamba_copies, + [0, 8, 15, 16, 19, 20, 23, 28, 31], + [0, 7, 8, 11, 12, 15, 16, 19], ), - ): - target = _attach_tensor_attrs( - torch.zeros(target_shape), - weight_loader=MethodType(lambda _owner: None, SimpleNamespace()), - ) - ext = _make_sparse_delta_extension( - mamba_name, - target, - SimpleNamespace( - tp_size=2, - intermediate_size=8, - groups_ssm_state_size=groups, - num_heads=4, - ), - ) - plan = ext._direct_sparse_delta_mamba2_plan( - {"name": mamba_name, "shape": (28, 2)}, - mamba_name, - {mamba_name: target}, - ) - _assert_sparse_plan(ext, plan, source_locations, expected_locations, values) + ) + for copies, source_locations, expected_rows in cases: + mapped = [ + VllmSparseDeltaApplier._map_copy( + torch.tensor(source_locations), + torch.ones(len(source_locations)), + copy, + )[0] + for copy in copies + ] + assert torch.cat(mapped).tolist() == expected_rows - for attrs, source_shape, source_locations, expected_locations, values in ( - ( - {"output_dim": 0}, - (6, 2), - [0, 1, 6, 7, 10, 11], - [0, 1, 4, 5], - [2, 3, 4, 5], - ), - ( - {"output_dim": 0, "input_dim": 1}, - (3, 4), - [0, 1, 2, 3, 6, 7, 10, 11], - [0, 1, 2, 3, 4, 5], - [2, 3, 4, 5, 6, 7], - ), - ): - target = _attach_tensor_attrs(torch.zeros(3, 2), **attrs, tp_size=2, tp_rank=1) - ext = _make_sparse_delta_extension("down_proj.weight", target, object()) - plan = ext._direct_sparse_delta_shard_plan( - {"name": "down_proj.weight", "shape": source_shape}, target + +@pytest.mark.vllm +def test_unknown_native_loader_fails_closed(monkeypatch) -> None: + _stub_fp8(monkeypatch) + model = _NativeLoaderModel(identity=torch.zeros(1)) + for name, error in (("unknown", "did not place"), ("transformed", "transform")): + payload = encode_sparse_infos( + [(name, torch.empty(1), torch.tensor([0]), torch.tensor([1.0]))], + empty_dtype=torch.float32, ) - _assert_sparse_plan(ext, plan, source_locations, expected_locations, values) + with pytest.raises(RuntimeError, match=error): + _applier(model)._apply_sparse_weight_deltas(payload[:2], payload[2]) @pytest.mark.vllm @pytest.mark.parametrize( ("initial", "expected_delta", "exact_mismatches", "mismatches"), - [ - (200.0, 4.0, 0, 0), - (2.0, 4.0000005, 1, 0), - (2.0, 5.0, 1, 1), - ], + [(200.0, 4.0, 0, 0), (2.0, 4.0000005, 1, 0), (2.0, 5.0, 1, 1)], ) -def test_sparse_delta_sample_verification_only_compares_applied_delta( +def test_sparse_delta_verification_compares_applied_delta( monkeypatch, initial: float, expected_delta: float, exact_mismatches: int, mismatches: int, ) -> None: - monkeypatch.setattr(torch.cuda, "is_available", lambda: False) - monkeypatch.setattr( - "nemo_rl.models.generation.vllm.quantization.fp8.is_fp8_model", - lambda _config: False, - ) + _stub_fp8(monkeypatch) target = torch.tensor([1.0, initial, 3.0, initial]) - model_runner = SimpleNamespace( - model=SimpleNamespace(), - vllm_config=SimpleNamespace( - model_config=SimpleNamespace(architectures=[]), - ), - ) - ext = VllmSparseDeltaApplier(model_runner, torch.device("cpu")) - ext._direct_sparse_delta_targets = {"weight": target} - ext._direct_sparse_delta_plan_cache = { - "weight": ext._make_sparse_delta_target_plan(target, (4,)) - } payload = encode_sparse_infos( [("weight", target, torch.tensor([1, 3]), torch.tensor([4.0, 4.0]))], empty_dtype=target.dtype, ) - metadata = payload[2] - metadata[0].update( + payload[2][0].update( verification_locations=[1, 3], verification_deltas=[expected_delta, expected_delta], ) + applier = _applier(_NativeLoaderModel(identity=target)) - ext._apply_sparse_weight_deltas(payload[:2], metadata) - result = ext.finish_sparse_delta_refit() + applier._apply_sparse_weight_deltas(payload[:2], payload[2]) + result = applier.finish_sparse_delta_refit() assert torch.equal(target, torch.tensor([1.0, initial + 4.0, 3.0, initial + 4.0])) assert result["verification_candidates"] == 2 assert result["verification_samples"] == 2 assert result["verification_exact_mismatches"] == 2 * exact_mismatches assert result["verification_mismatches"] == 2 * mismatches - rounded_difference = float((torch.tensor(expected_delta) - 4).abs()) - assert result["verification_max_abs"] == rounded_difference diff --git a/tests/unit/models/generation/test_vllm_sparse_refit.py b/tests/unit/models/generation/test_vllm_sparse_refit.py index 21a821643d8..a3b9e0f2fbd 100644 --- a/tests/unit/models/generation/test_vllm_sparse_refit.py +++ b/tests/unit/models/generation/test_vllm_sparse_refit.py @@ -95,7 +95,7 @@ def apply(payloads: tuple[bytes, ...]) -> dict[str, Any]: assert response["payloads"] == 5 assert response["batches"] == 2 assert sum(result.get("receiver_total_s", 0.0) for result in responses) == 5.0 - receiver.llm.collective_rpc.assert_called_once_with( + receiver._worker.llm.collective_rpc.assert_called_once_with( "finish_sparse_delta_refit", args=() ) @@ -209,7 +209,7 @@ def collective_rpc(method: str, args: tuple[str, ...]) -> list[dict[str, Any]]: assert staged_payloads == [b"0", b"1", b"2"] assert not list(tmp_path.iterdir()) - receiver.llm.collective_rpc.assert_called_once() + receiver._worker.llm.collective_rpc.assert_called_once() assert response == {"ok": True, "receiver_total_s": 1.0, "payloads": 3} @@ -236,7 +236,8 @@ def collective_rpc(method: str, args: tuple[str, ...]) -> list[Any]: receiver.update_weights_from_serialized_sparse_payloads((b"0", b"1")) assert [ - entry.args[0] for entry in receiver.llm.collective_rpc.call_args_list + entry.args[0] + for entry in receiver._worker.llm.collective_rpc.call_args_list ] == [ "update_weights_from_sparse_payload_files", "synchronize_device", @@ -244,7 +245,7 @@ def collective_rpc(method: str, args: tuple[str, ...]) -> list[Any]: assert not list(tmp_path.iterdir()) -def test_sparse_refit_batch_falls_back_across_nodes() -> None: +def test_sparse_refit_batch_uses_one_collective_rpc_across_nodes() -> None: with _sparse_refit_receiver() as receiver: receiver._refit_workers_share_node = False receiver._worker.llm = MagicMock( @@ -257,11 +258,11 @@ def test_sparse_refit_batch_falls_back_across_nodes() -> None: (b"0", b"1", b"2") ) - assert receiver.llm.collective_rpc.call_args_list == [ - call("update_weights_from_serialized_sparse_payload", args=(payload,)) - for payload in (b"0", b"1", b"2") - ] - assert response == {"ok": True, "receiver_total_s": 3.0, "payloads": 3} + receiver._worker.llm.collective_rpc.assert_called_once_with( + "update_weights_from_serialized_sparse_payload", + args=(b"0", b"1", b"2"), + ) + assert response == {"ok": True, "receiver_total_s": 1.0, "payloads": 3} @pytest.mark.asyncio diff --git a/tests/unit/utils/test_weight_transfer_remote_sparse.py b/tests/unit/utils/test_weight_transfer_remote_sparse.py index cf48be816ec..3add0c4ff31 100644 --- a/tests/unit/utils/test_weight_transfer_remote_sparse.py +++ b/tests/unit/utils/test_weight_transfer_remote_sparse.py @@ -39,11 +39,33 @@ ) +class _SparsePipelineTracker: + sparse_bucket_size_bytes = 1 + + @staticmethod + def prepare_sparse_delta_payload(chunk): + return (chunk, torch.ones(1), [1]), 1, 1 + + +def _stream_sparse_test_payloads(tensors, send_payload): + return weight_transfer_remote_sparse.stream_sparse_delta_payloads( + tensors, + delta_tracker=_SparsePipelineTracker(), + transport="zmq", + send_payload=send_payload, + transfer_workers=1, + shard_rank=0, + shard_count=1, + ) + + +def _delta_tracker() -> DeltaCompressionTracker: + return DeltaCompressionTracker({"dtype": "bf16", "sparse_bucket_size_bytes": 1024}) + + def test_delta_tracker_commits_only_successful_syncs(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") - tracker = DeltaCompressionTracker( - {"dtype": "bf16", "sparse_bucket_size_bytes": 1024} - ) + tracker = _delta_tracker() tensor = torch.tensor([1.0, 2.0, 3.0]) tracker.snapshot_baseline([("weight", tensor)]) tensor[1] += 4 @@ -58,9 +80,7 @@ def test_delta_tracker_commits_only_successful_syncs(monkeypatch) -> None: def test_delta_tracker_emits_bounded_delta_samples(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") monkeypatch.setenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "2") - tracker = DeltaCompressionTracker( - {"dtype": "bf16", "sparse_bucket_size_bytes": 1024} - ) + tracker = _delta_tracker() tensor = torch.tensor([1.0, 2.0, 3.0, 4.0]) tracker.snapshot_baseline([("weight", tensor)]) tensor[[1, 3]] += 1 @@ -76,9 +96,7 @@ def test_delta_tracker_emits_bounded_delta_samples(monkeypatch) -> None: def test_delta_tracker_commits_quantized_receiver_baseline(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") - tracker = DeltaCompressionTracker( - {"dtype": "bf16", "sparse_bucket_size_bytes": 1024} - ) + tracker = _delta_tracker() tensor = torch.tensor([1.0]) tracker.snapshot_baseline([("weight", tensor)]) tensor.add_(0.001) @@ -141,13 +159,6 @@ def test_sparse_export_finishes_before_blocked_transfers(monkeypatch) -> None: release_transfers = threading.Event() result = [] - class Tracker: - sparse_bucket_size_bytes = 1 - - @staticmethod - def prepare_sparse_delta_payload(chunk): - return (chunk, torch.ones(1), [1]), 1, 1 - def tensors(): for index in range(4): yield f"weight-{index}", torch.ones(1) @@ -158,17 +169,7 @@ def send_payload(_body, _payload_index): return {"receiver": {}} def run(): - result.append( - weight_transfer_remote_sparse.stream_sparse_delta_payloads( - tensors(), - delta_tracker=Tracker(), - transport="zmq", - send_payload=send_payload, - transfer_workers=1, - shard_rank=0, - shard_count=1, - ) - ) + result.append(_stream_sparse_test_payloads(tensors(), send_payload)) thread = threading.Thread(target=run) thread.start() @@ -187,13 +188,6 @@ def test_sparse_export_finishes_before_transfer_error(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "1") exported = [] - class Tracker: - sparse_bucket_size_bytes = 1 - - @staticmethod - def prepare_sparse_delta_payload(chunk): - return (chunk, torch.ones(1), [1]), 1, 1 - def tensors(): for index in range(4): exported.append(index) @@ -203,15 +197,7 @@ def fail_transfer(_body, _payload_index): raise RuntimeError("transfer failed") with pytest.raises(RuntimeError, match="transfer failed"): - weight_transfer_remote_sparse.stream_sparse_delta_payloads( - tensors(), - delta_tracker=Tracker(), - transport="zmq", - send_payload=fail_transfer, - transfer_workers=1, - shard_rank=0, - shard_count=1, - ) + _stream_sparse_test_payloads(tensors(), fail_transfer) assert exported == list(range(4)) @@ -483,47 +469,21 @@ def test_zmq_server_rejects_malformed_messages() -> None: "checksum": sparse_payload_checksum(body), } - with pytest.raises(ValueError, match="Expected 4"): - server._parse_data_message([]) - with pytest.raises(ValueError, match="Unsupported ZeroMQ sparse refit message"): - server._parse_data_message( - [b"id", b"OTHER", json.dumps(metadata).encode(), body] - ) - with pytest.raises(ValueError, match="protocol"): - server._parse_data_message( - [ - b"id", - b"DATA", - json.dumps({**metadata, "protocol": "other"}).encode(), - body, - ] - ) - with pytest.raises(PermissionError, match="authentication"): - server._parse_data_message( - [ - b"id", - b"DATA", - json.dumps({**metadata, "api_key": "wrong"}).encode(), - body, - ] - ) - with pytest.raises(ValueError, match="identity"): - server._parse_data_message( - [b"id", b"DATA", json.dumps({**metadata, "transfer_id": ""}).encode(), body] - ) - with pytest.raises(ValueError, match="checksum mismatch"): - server._parse_data_message( - [ - b"id", - b"DATA", - json.dumps({**metadata, "checksum": "wrong"}).encode(), - body, - ] - ) - - assert server._parse_data_message( - [b"id", b"DATA", json.dumps(metadata).encode(), body] - )[1] == ("transfer", 0, 1) + def frames(kind: bytes = b"DATA", **updates: object) -> list[bytes]: + return [b"id", kind, json.dumps({**metadata, **updates}).encode(), body] + + for message, error, match in ( + ([], ValueError, "Expected 4"), + (frames(b"OTHER"), ValueError, "Unsupported ZeroMQ sparse refit message"), + (frames(protocol="other"), ValueError, "protocol"), + (frames(api_key="wrong"), PermissionError, "authentication"), + (frames(transfer_id=""), ValueError, "identity"), + (frames(checksum="wrong"), ValueError, "checksum mismatch"), + ): + with pytest.raises(error, match=match): + server._parse_data_message(message) + + assert server._parse_data_message(frames())[1] == ("transfer", 0, 1) def _receiver_server(received): 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 6b5683e90a5..5107241abb3 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 @@ -94,8 +94,8 @@ def test_shutdown_cancels_pending_work_and_stops_zmq(self, mock_ray): run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], ) assert sync.is_stale - assert sync._baseline_init_refs is None - assert sync._baseline_commit_refs is None + assert sync._baseline_init_refs == [] + assert sync._baseline_commit_refs == [] assert sync._refit_urls == [] assert sync._targets == [] diff --git a/tools/refit_bandwidth_calculator.py b/tools/refit_bandwidth_calculator.py index be82173a0b0..b4bcb0bc5e8 100644 --- a/tools/refit_bandwidth_calculator.py +++ b/tools/refit_bandwidth_calculator.py @@ -12,64 +12,41 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Estimate when measured sparse refit beats NCCL over Ethernet. - -This is a benchmark-specific estimator. Sparse latency is fitted from the July -2026 S3 and ZeroMQ benchmarks, and NCCL latency is interpolated from measured -H100 reshard results on 400 Gbps/rank InfiniBand using hierarchical API. -``--candidate-ethernet-gbps`` projects that NCCL reference onto a raw -per-rank Ethernet rate. -""" +"""Compare measured sparse refit with NCCL projected from H100 IB to Ethernet.""" import argparse import json import math +from bisect import bisect_right from dataclasses import asdict, dataclass -from itertools import pairwise from typing import Literal Transport = Literal["s3", "zmq"] Compression = Literal["raw", "zstd"] _REFERENCE_IB_GBPS = 400.0 -_CALIBRATED_DENSITIES = (3.0, 5.0) -_CALIBRATED_MODEL_SIZE_GB = (63.2, 1121.0) - - -@dataclass(frozen=True) -class _SparseFit: - intercept_s: float - seconds_per_tb: float - - -# T(S) = intercept_s + seconds_per_tb * S for decimal TB of indexed BF16. -_SPARSE_FITS: dict[tuple[Transport, Compression, float], _SparseFit] = { - ("s3", "raw", 3.0): _SparseFit(2.370, 116.138), - ("s3", "zstd", 3.0): _SparseFit(0.000, 84.746), - ("s3", "raw", 5.0): _SparseFit(2.646, 183.431), - ("s3", "zstd", 5.0): _SparseFit(0.000, 134.041), - ("zmq", "raw", 3.0): _SparseFit(0.000, 264.753), - ("zmq", "zstd", 3.0): _SparseFit(6.517, 73.621), - ("zmq", "raw", 5.0): _SparseFit(0.000, 428.008), - ("zmq", "zstd", 5.0): _SparseFit(8.165, 157.339), +_DENSITIES = (3.0, 5.0) +_SPARSE_SIZE_RANGE_GB = (63.2, 1121.0) +# Each pair is (fixed seconds, seconds per 1,000 GB) at 3% and 5%. +_SPARSE_LATENCY_FITS: dict[ + tuple[Transport, Compression], tuple[tuple[float, float], tuple[float, float]] +] = { + ("s3", "raw"): ((2.370, 116.138), (2.646, 183.431)), + ("s3", "zstd"): ((0.000, 84.746), (0.000, 134.041)), + ("zmq", "raw"): ((0.000, 264.753), (0.000, 428.008)), + ("zmq", "zstd"): ((6.517, 73.621), (8.165, 157.339)), } - -# (indexed model GB, low seconds, high seconds), with generation EP enabled. _NCCL_ANCHORS = ( (63.2, 0.84, 1.60), (247.2, 1.46, 1.74), (470.2, 2.31, 2.73), (1342.0, 3.27, 3.46), ) - -# Approximate serialized bytes / changed BF16 bytes. _WIRE_MULTIPLIER: dict[Compression, float] = {"raw": 2.0, "zstd": 0.74} @dataclass(frozen=True) class Estimate: - """One sparse transport compared with the NCCL reference.""" - transport: Transport compression: Compression model_size_gb: float @@ -84,8 +61,6 @@ class Estimate: break_even_ethernet_low_gbps: float break_even_ethernet_high_gbps: float candidate_winner: str | None - model_size_extrapolated: bool - density_extrapolated: bool def predict_sparse_seconds( @@ -95,50 +70,24 @@ def predict_sparse_seconds( transport: Transport, compression: Compression, ) -> float: - """Interpolate or extrapolate sparse latency from the 3% and 5% fits.""" + """Evaluate the campaign fit at any positive changed density.""" if model_size_gb <= 0 or changed_pct <= 0: raise ValueError("model_size_gb and changed_pct must be positive") - size_tb = model_size_gb / 1000.0 - density_low, density_high = _CALIBRATED_DENSITIES - latency_low = _fit_latency(size_tb, transport, compression, density_low) - latency_high = _fit_latency(size_tb, transport, compression, density_high) - exponent = math.log(latency_high / latency_low) / math.log( - density_high / density_low - ) - return latency_low * (changed_pct / density_low) ** exponent + (a3, b3), (a5, b5) = _SPARSE_LATENCY_FITS[(transport, compression)] + latency_3, latency_5 = a3 + b3 * size_tb, a5 + b5 * size_tb + exponent = math.log(latency_5 / latency_3) / math.log(5.0 / 3.0) + return latency_3 * (changed_pct / 3.0) ** exponent -def _fit_latency( - size_tb: float, - transport: Transport, - compression: Compression, - changed_pct: float, -) -> float: - fit = _SPARSE_FITS[(transport, compression, changed_pct)] - return fit.intercept_s + fit.seconds_per_tb * size_tb - - -def _nccl_reference(model_size_gb: float) -> tuple[float, float, bool]: - if model_size_gb <= 0: - raise ValueError("model_size_gb must be positive") - - extrapolated = not _NCCL_ANCHORS[0][0] <= model_size_gb <= _NCCL_ANCHORS[-1][0] - if model_size_gb <= _NCCL_ANCHORS[0][0]: - left, right = _NCCL_ANCHORS[:2] - elif model_size_gb >= _NCCL_ANCHORS[-1][0]: - left, right = _NCCL_ANCHORS[-2:] - else: - left, right = _NCCL_ANCHORS[:2] - for candidate_left, candidate_right in pairwise(_NCCL_ANCHORS): - if candidate_left[0] <= model_size_gb <= candidate_right[0]: - left, right = candidate_left, candidate_right - break - +def _nccl_reference(model_size_gb: float) -> tuple[float, float]: + sizes = tuple(anchor[0] for anchor in _NCCL_ANCHORS) + index = min(max(bisect_right(sizes, model_size_gb) - 1, 0), len(sizes) - 2) + left, right = _NCCL_ANCHORS[index : index + 2] position = math.log(model_size_gb / left[0]) / math.log(right[0] / left[0]) - low_s = left[1] + position * (right[1] - left[1]) - high_s = left[2] + position * (right[2] - left[2]) - return max(0.001, low_s), max(0.001, high_s), extrapolated + low = left[1] + position * (right[1] - left[1]) + high = left[2] + position * (right[2] - left[2]) + return max(0.001, low), max(0.001, high) def estimate( @@ -152,93 +101,58 @@ def estimate( """Estimate sparse latency and the NCCL-over-Ethernet crossover.""" if candidate_ethernet_gbps is not None and candidate_ethernet_gbps <= 0: raise ValueError("candidate_ethernet_gbps must be positive") - sparse_seconds = predict_sparse_seconds( model_size_gb, changed_pct, transport=transport, compression=compression, ) - nccl_low_s, nccl_high_s, model_size_extrapolated = _nccl_reference(model_size_gb) - break_even_low = _REFERENCE_IB_GBPS * nccl_low_s / sparse_seconds - break_even_high = _REFERENCE_IB_GBPS * nccl_high_s / sparse_seconds + nccl_low, nccl_high = _nccl_reference(model_size_gb) + crossover_low = _REFERENCE_IB_GBPS * nccl_low / sparse_seconds + crossover_high = _REFERENCE_IB_GBPS * nccl_high / sparse_seconds - candidate_low = candidate_high = None + projected_low = projected_high = None winner = None if candidate_ethernet_gbps is not None: scale = _REFERENCE_IB_GBPS / candidate_ethernet_gbps - candidate_low = nccl_low_s * scale - candidate_high = nccl_high_s * scale - if sparse_seconds < candidate_low and not math.isclose( - sparse_seconds, candidate_low + projected_low, projected_high = nccl_low * scale, nccl_high * scale + if sparse_seconds < projected_low and not math.isclose( + sparse_seconds, projected_low ): winner = transport - elif sparse_seconds > candidate_high and not math.isclose( - sparse_seconds, candidate_high + elif sparse_seconds > projected_high and not math.isclose( + sparse_seconds, projected_high ): winner = "nccl" else: winner = "depends" return Estimate( - transport=transport, - compression=compression, - model_size_gb=model_size_gb, - changed_pct=changed_pct, - sparse_seconds=sparse_seconds, - approximate_wire_gb=( - model_size_gb * changed_pct / 100.0 * _WIRE_MULTIPLIER[compression] - ), - nccl_ib_low_s=nccl_low_s, - nccl_ib_high_s=nccl_high_s, - candidate_ethernet_gbps=candidate_ethernet_gbps, - nccl_ethernet_low_s=candidate_low, - nccl_ethernet_high_s=candidate_high, - break_even_ethernet_low_gbps=break_even_low, - break_even_ethernet_high_gbps=break_even_high, - candidate_winner=winner, - model_size_extrapolated=( - model_size_extrapolated - or not _CALIBRATED_MODEL_SIZE_GB[0] - <= model_size_gb - <= _CALIBRATED_MODEL_SIZE_GB[1] - ), - density_extrapolated=( - not _CALIBRATED_DENSITIES[0] <= changed_pct <= _CALIBRATED_DENSITIES[1] - ), + transport, + compression, + model_size_gb, + changed_pct, + sparse_seconds, + model_size_gb * changed_pct / 100 * _WIRE_MULTIPLIER[compression], + nccl_low, + nccl_high, + candidate_ethernet_gbps, + projected_low, + projected_high, + crossover_low, + crossover_high, + winner, ) -def _positive_float(value: str) -> float: +def _positive(value: str) -> float: number = float(value) if number <= 0: raise argparse.ArgumentTypeError("must be positive") return number -def parse_args() -> argparse.Namespace: - """Parse command-line arguments.""" - parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--model-size-gb", type=_positive_float, required=True) - parser.add_argument( - "--changed-pct", - "--sparsity-pct", - type=_positive_float, - required=True, - help="Any positive percentage of changed weight elements.", - ) - parser.add_argument("--transport", choices=("all", "s3", "zmq"), default="all") - parser.add_argument("--compression", choices=("raw", "zstd"), default="zstd") - parser.add_argument( - "--candidate-ethernet-gbps", - type=_positive_float, - help="Raw per-rank Ethernet rate used by the projected NCCL refit.", - ) - parser.add_argument("--json", action="store_true") - return parser.parse_args() - - -def _seconds_range(low: float, high: float) -> str: +def _seconds(low: float, high: float) -> str: return f"{low:.3f}-{high:.3f}s" @@ -249,53 +163,51 @@ def _print_results(results: list[Estimate]) -> None: f"changed: {first.changed_pct:g}%; compression: {first.compression}" ) print( - "Measured NCCL on 400 Gbps/rank IB: " - f"{_seconds_range(first.nccl_ib_low_s, first.nccl_ib_high_s)}" + "Measured NCCL on 400 Gbps/rank H100 IB: " + f"{_seconds(first.nccl_ib_low_s, first.nccl_ib_high_s)}" ) - if first.candidate_ethernet_gbps is not None: assert first.nccl_ethernet_low_s is not None assert first.nccl_ethernet_high_s is not None print( f"Projected NCCL on {first.candidate_ethernet_gbps:g} Gbps/rank " - f"Ethernet: {_seconds_range(first.nccl_ethernet_low_s, first.nccl_ethernet_high_s)}" + f"Ethernet: {_seconds(first.nccl_ethernet_low_s, first.nccl_ethernet_high_s)}" ) - print() - winner_header = "Winner@candidate" if first.candidate_ethernet_gbps else "" print( - f"{'Path':<6} {'Sparse':>10} {'Wire':>10} " - f"{'Ethernet crossover':>26} {winner_header:>18}" + "\nPath Sparse Wire Ethernet crossover Winner@candidate" ) for result in results: crossover = ( f"{result.break_even_ethernet_low_gbps:.2f}-" f"{result.break_even_ethernet_high_gbps:.2f} Gbps/rank" ) - winner = result.candidate_winner or "" print( f"{result.transport.upper():<6} {result.sparse_seconds:>9.3f}s " - f"{result.approximate_wire_gb:>8.3f} GB {crossover:>26} {winner:>18}" + f"{result.approximate_wire_gb:>8.3f} GB {crossover:>26} " + f"{result.candidate_winner or '':>18}" ) - - print() print( - "Below the lower crossover, sparse refit beats the full NCCL envelope; " - "above the upper crossover, NCCL wins." + "\nBelow the lower crossover sparse wins across the NCCL envelope; " + "above the upper crossover NCCL wins." ) - print( - "Projection: T_ethernet = T_H100_IB * 400 / candidate_gbps, with all " - "bandwidth values expressed per rank." - ) - if first.model_size_extrapolated: - print("Note: model size is outside the measured calibration range.") - if first.density_extrapolated: + if not _SPARSE_SIZE_RANGE_GB[0] <= first.model_size_gb <= _SPARSE_SIZE_RANGE_GB[1]: + print("Note: model size is outside the measured sparse calibration range.") + if not _DENSITIES[0] <= first.changed_pct <= _DENSITIES[1]: print("Note: changed density is extrapolated from measured 3% and 5% arms.") def main() -> None: - """Run the estimator.""" - args = parse_args() + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--model-size-gb", type=_positive, required=True) + parser.add_argument( + "--changed-pct", "--sparsity-pct", type=_positive, required=True + ) + parser.add_argument("--transport", choices=("all", "s3", "zmq"), default="all") + parser.add_argument("--compression", choices=("raw", "zstd"), default="zstd") + parser.add_argument("--candidate-ethernet-gbps", type=_positive) + parser.add_argument("--json", action="store_true") + args = parser.parse_args() transports: tuple[Transport, ...] = ( ("s3", "zmq") if args.transport == "all" else (args.transport,) ) From 9d386c3e58883444f33013ac80e8503b896b2c45 Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Sat, 11 Jul 2026 14:18:32 -0700 Subject: [PATCH 07/18] only keep xor and overwrite Signed-off-by: Hollow Man --- docs/design-docs/sparse-delta-refit.md | 49 +-- nemo_rl/models/generation/vllm/config.py | 2 +- .../generation/vllm/vllm_sparse_delta.py | 225 ++++++++---- nemo_rl/utils/weight_transfer_sparse_codec.py | 170 ++++++--- .../generation/test_vllm_sparse_delta.py | 338 ++++++++++++++++-- .../test_weight_transfer_remote_sparse.py | 124 ++++++- 6 files changed, 739 insertions(+), 169 deletions(-) diff --git a/docs/design-docs/sparse-delta-refit.md b/docs/design-docs/sparse-delta-refit.md index 62932993f01..d0022be6e9e 100644 --- a/docs/design-docs/sparse-delta-refit.md +++ b/docs/design-docs/sparse-delta-refit.md @@ -3,8 +3,8 @@ Remote sparse-delta refit updates non-colocated vLLM workers without sending a full checkpoint after every optimizer step. Megatron workers export Hugging Face (HF) weights, compare them with a sharded CPU baseline, and send changed -locations and deltas through S3 or ZeroMQ. vLLM maps those HF coordinates into -its local TP and EP layouts and applies the updates in place. +locations and byte-encoded values through S3 or ZeroMQ. vLLM maps those HF +coordinates into its local TP and EP layouts and applies the updates in place. The feature is opt-in. Its synchronizer, codec, transports, receiver queue, and placement engine are separate from existing NCCL, CUDA IPC, and packed refit @@ -91,12 +91,12 @@ not remove Bridge export or the full CPU comparison. Low changed density mainly reduces encoded and transferred bytes. For each assigned chunk, `DeltaCompressionTracker` finds changed flat -locations and encodes deltas in the configured dtype. The pending baseline is -updated to the value the receiver will hold after wire-dtype rounding: - -```text -expected = previous_baseline + delta.to(baseline_dtype) -``` +locations through an integer view with the same element width. It encodes +either the absolute new bits (`overwrite`) or new bits XOR baseline bits +(`xor`). Both encodings are dtype-blind and preserve FP8 bit patterns. +`overwrite` is idempotent and recommended. `xor` can compress better, but it +requires an exact same-dtype receiver baseline and exactly-once application. +The pending source baseline always records the exact new source bits. The producer overlaps export, encoding, `torch.save` serialization, zstd level 1 compression, and transfer with fixed-size executors. Source baselines do not @@ -128,33 +128,37 @@ updates in background CPU threads. > **Failure boundary:** source baseline commit is transactional, but receiver > updates are in place and are not rolled back. If a transfer fails after a > receiver accepts any payload, reload that receiver from a known-good weight -> version before retrying. +> version before retrying. This is mandatory for `xor`, because replaying an +> already-applied XOR reverts those bits. Replaying `overwrite` is safe. ## Payload and placement Each serialized payload is: ```text -(packed_location_bytes, packed_delta_values, tensor_metadata) +(packed_location_bytes, packed_value_groups, tensor_metadata) ``` Contiguous locations use a range encoding. Other sorted locations are delta-encoded into the smallest lossless unsigned width among 16, 32, and 64 bits. Metadata carries the HF name and shape, value offsets, location encoding, -and optional verification samples. +the `xor` or `overwrite` operation, and optional verification samples. HF coordinates are the canonical wire format because Megatron Bridge already defines the training-to-HF mapping while vLLM owns a different packed and sharded layout. On first use, the receiver runs vLLM's native `load_weights()` against metadata-only tensors while a PyTorch dispatch mode records the source and destination views of each `copy_`. It caches those mappings and applies -later sparse deltas directly with `index_add_`, without materializing a dense HF -tensor or duplicating QKV, MoE, Mamba, or TP placement rules. +later sparse values directly with bitwise XOR or `index_copy_`, without +materializing a dense HF tensor or duplicating QKV, MoE, Mamba, or TP placement +rules. -The tracer accepts affine tensor views and the Mamba `A_log` transform. An -element-expanding copy, unknown transform, or unplaced -non-expert tensor fails before any payload update. There is no dense fallback -for an unknown layout. +The tracer accepts affine tensor views. Absolute overwrite also supports the +Mamba `A_log` transform and source-to-target dtype casts. XOR rejects either +case because target bits no longer match the source representation. XOR also +rejects overlapping target mappings. An element-expanding copy, unknown +transform, or unplaced non-expert tensor fails before any payload update. There +is no dense fallback for an unknown layout. ## Configuration @@ -166,7 +170,7 @@ policy: backend: vllm refit_transport: vllm_s3_sparse # or vllm_zmq_sparse delta_compression: - dtype: bf16 + encoding: overwrite # xor requires an exact baseline and exactly-once apply sparse_bucket_size_bytes: 268435456 colocated: enabled: false @@ -233,18 +237,19 @@ Keep transport changes behind the shared `stream_sparse_delta_payloads()` pipeline. A transport should provide payload delivery and timing only; it must not duplicate the baseline tracker, codec, receiver queue, or placement logic. Retries must preserve payload identity and bytes, fan out to every required -replica, and require a successful global flush before baseline commit. +replica, and require a successful global flush before baseline commit. Never +retry XOR after an uncertain or partial receiver apply. Do not add model-specific placement math. New layouts should work through their native vLLM weight loader; extend the tracer only for a general loader operation and fail closed for transformed or broadcasting copies. Tests must invoke the real vLLM loader at nonzero TP ranks and cover replicated KV heads, packed columns, local and remote experts, segmented views, contiguous ranges, and -explicit locations. Incorrect in-range `index_add_` locations silently corrupt -weights, so assert exact mapped indices and values. +explicit locations. Incorrect in-range XOR or overwrite locations silently +corrupt weights, so assert exact mapped indices and values. Codec changes must update encoder and decoder together, preserve 64-bit-safe -locations, and retain wire-dtype rounding in pending baseline updates. Receiver +locations, and commit exact source bits only after global success. Receiver changes must preserve FIFO application, bounded memory, deferred-error propagation, flush, CUDA synchronization, and clean shutdown. diff --git a/nemo_rl/models/generation/vllm/config.py b/nemo_rl/models/generation/vllm/config.py index 62329563b34..23e6885b457 100644 --- a/nemo_rl/models/generation/vllm/config.py +++ b/nemo_rl/models/generation/vllm/config.py @@ -59,7 +59,7 @@ class VllmSpecificArgs(TypedDict): class VllmDeltaCompressionConfig(TypedDict): - dtype: Literal["fp16", "float16", "bf16", "bfloat16", "fp32", "float32"] # fmt: skip + encoding: Literal["xor", "overwrite"] sparse_bucket_size_bytes: int diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py index e0569c8b6fc..35432bd8d4d 100644 --- a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py +++ b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py @@ -138,9 +138,7 @@ def __init__(self, model_runner: Any, device: torch.device) -> None: self.model_runner = model_runner self._cuda_device_index = device.index self._plan_cache: dict[str, _SparseDeltaTargetPlan] = {} - self._verification: list[ - tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] - ] = [] + self._verification: list[tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = [] self._verification_candidates = 0 def _compile_plans(self, metadata: list[dict[str, Any]]) -> None: @@ -155,7 +153,11 @@ def _compile_plans(self, metadata: list[dict[str, Any]]) -> None: model = self.model_runner.model targets = list(model.parameters()) + list(model.buffers()) sources = { - name: torch.empty(tuple(item["shape"]), device="meta") + name: torch.empty( + tuple(item["shape"]), + dtype=sparse_codec.dtype_from_name(str(item["dtype"])), + device="meta", + ) for name, item in missing.items() } tracer = _SparseLoadTracer(targets, sources) @@ -227,96 +229,178 @@ def _record_verification( self, item: dict[str, Any], plan: _SparseDeltaTargetPlan, + operation: sparse_codec.SparseOperation, + source_dtype: torch.dtype, ) -> None: sample_locations = item.get("verification_locations", []) self._verification_candidates += len(sample_locations) - if not sample_locations or plan.log_delta_transform or not plan.copies: + if not sample_locations or not plan.copies: return target = plan.copies[0].target locations = torch.tensor(sample_locations, device=target.device) + value_dtype = sparse_codec.integer_dtype_for_element_size(source_dtype.itemsize) values = torch.tensor( - item["verification_deltas"], device=target.device, dtype=target.dtype + item["verification_values"], device=target.device, dtype=value_dtype ) for copy in plan.copies: mapped, selected = self._map_copy(locations, values, copy) if not mapped.numel(): continue - before = copy.target.data.view(-1).index_select(0, mapped) - expected = (before + selected).float() - before.float() - self._verification.append((copy.target, mapped, before.float(), expected)) + if operation == "xor": + target_bits = self._integer_flat(copy.target) + expected = target_bits.index_select(0, mapped).bitwise_xor(selected) + else: + _, replacement = self._overwrite_target_values( + copy.target, + selected, + source_dtype, + log_transform=plan.log_delta_transform, + ) + expected = replacement.contiguous().view( + sparse_codec.integer_dtype_for_element_size( + copy.target.element_size() + ) + ) + self._verification.append((copy.target, mapped, expected)) + + @staticmethod + def _integer_flat(target: torch.Tensor) -> torch.Tensor: + dtype = sparse_codec.integer_dtype_for_element_size(target.element_size()) + return target.data.view(dtype).view(-1) + + @staticmethod + def _xor_target_mappings_overlap(plan: _SparseDeltaTargetPlan) -> bool: + spans: dict[int, list[tuple[int, int]]] = defaultdict(list) + for copy in plan.copies: + origin = int(copy.target.storage_offset()) + copy.target_offset + extents = [ + (size - 1) * stride + for size, stride in zip(copy.shape, copy.target_strides, strict=True) + ] + start = origin + sum(min(0, extent) for extent in extents) + end = origin + sum(max(0, extent) for extent in extents) + target_spans = spans[_storage_key(copy.target)] + if any( + start <= other_end and other_start <= end + for other_start, other_end in target_spans + ): + return True + target_spans.append((start, end)) + return False + + @classmethod + def _overwrite_target_values( + cls, + target: torch.Tensor, + values: torch.Tensor, + source_dtype: torch.dtype, + *, + log_transform: bool, + ) -> tuple[torch.Tensor, torch.Tensor]: + if not log_transform and target.dtype == source_dtype: + return cls._integer_flat(target), values + source_values = values.contiguous().view(source_dtype) + replacement = ( + -source_values.float().exp().to(target.dtype) + if log_transform + else source_values.to(target.dtype) + ) + return target.data.view(-1), replacement def _apply_item( self, item: dict[str, Any], plan: _SparseDeltaTargetPlan, raw_locations: torch.Tensor, - raw_values: torch.Tensor, + raw_value_groups: tuple[torch.Tensor, ...], ) -> None: - self._record_verification(item, plan) + operation = sparse_codec.sparse_operation(item["operation"]) + source_dtype = sparse_codec.dtype_from_name(str(item["dtype"])) if not plan.copies: + self._record_verification(item, plan, operation, source_dtype) return first_target = plan.copies[0].target + if operation == "xor" and plan.log_delta_transform: + raise RuntimeError(f"XOR cannot apply transformed weight {item['name']!r}.") + if operation == "xor" and any( + copy.target.dtype != source_dtype for copy in plan.copies + ): + raise RuntimeError( + f"XOR source and target dtypes differ for {item['name']!r}." + ) + if operation == "xor" and self._xor_target_mappings_overlap(plan): + raise RuntimeError(f"XOR target mappings overlap for {item['name']!r}.") value_start, value_end = int(item["value_start"]), int(item["value_end"]) - values = raw_values[value_start:value_end].to( - device=first_target.device, dtype=first_target.dtype, non_blocking=True + values = raw_value_groups[int(item["value_group"])][value_start:value_end] + expected_dtype = sparse_codec.integer_dtype_for_element_size( + source_dtype.itemsize ) + if values.dtype != expected_dtype: + raise RuntimeError( + f"Sparse values have the wrong dtype for {item['name']!r}." + ) + values = values.to(device=first_target.device, non_blocking=True) + self._record_verification(item, plan, operation, source_dtype) if plan.identity and item["index_encoding"] == "range": - first_target.data.view(-1).narrow( - 0, int(item["range_start"]), value_end - value_start - ).add_(values) + if operation == "xor": + target = self._integer_flat(first_target) + target.narrow( + 0, int(item["range_start"]), value_end - value_start + ).bitwise_xor_(values) + else: + target, replacement = self._overwrite_target_values( + first_target, values, source_dtype, log_transform=False + ) + target.narrow( + 0, int(item["range_start"]), value_end - value_start + ).copy_(replacement) return locations = sparse_codec.sparse_locations_for_item( item, raw_locations, device=first_target.device ) if plan.identity: - first_target.data.view(-1).index_add_(0, locations, values) + if operation == "xor": + target = self._integer_flat(first_target) + current = target.index_select(0, locations) + target.index_copy_(0, locations, current.bitwise_xor(values)) + else: + target, replacement = self._overwrite_target_values( + first_target, values, source_dtype, log_transform=False + ) + target.index_copy_(0, locations, replacement) return - grouped: dict[ - int, tuple[torch.Tensor, list[torch.Tensor], list[torch.Tensor]] - ] = {} for copy in plan.copies: mapped, selected = self._map_copy(locations, values, copy) - if mapped.numel(): - _, mapped_parts, value_parts = grouped.setdefault( - id(copy.target), (copy.target, [], []) - ) - mapped_parts.append(mapped) - value_parts.append(selected) - for target, mapped_parts, value_parts in grouped.values(): - mapped = ( - torch.cat(mapped_parts) if len(mapped_parts) > 1 else mapped_parts[0] - ) - selected = ( - torch.cat(value_parts) if len(value_parts) > 1 else value_parts[0] - ) - target_flat = target.data.view(-1) - if plan.log_delta_transform: - current = target_flat.index_select(0, mapped) - target_flat.index_copy_( - 0, mapped, current * selected.float().exp().to(current.dtype) - ) + if not mapped.numel(): + continue + if operation == "xor": + target = self._integer_flat(copy.target) + current = target.index_select(0, mapped) + target.index_copy_(0, mapped, current.bitwise_xor(selected)) else: - target_flat.index_add_(0, mapped, selected) + target, replacement = self._overwrite_target_values( + copy.target, + selected, + source_dtype, + log_transform=plan.log_delta_transform, + ) + target.index_copy_(0, mapped, replacement) def _apply_sparse_weight_deltas( self, - payload_tensors: tuple[torch.Tensor, torch.Tensor], + payload_tensors: tuple[torch.Tensor, tuple[torch.Tensor, ...]], metadata: list[dict[str, Any]], ) -> None: - from nemo_rl.models.generation.vllm.quantization import fp8 - - if fp8.is_fp8_model(self.model_runner.vllm_config): - raise RuntimeError( - "Direct sparse delta refit does not support FP8 weights." - ) - self._compile_plans(metadata) - raw_locations, raw_values = payload_tensors + raw_locations, raw_value_groups = payload_tensors with torch.no_grad(): for item in metadata: self._apply_item( - item, self._plan_cache[str(item["name"])], raw_locations, raw_values + item, + self._plan_cache[str(item["name"])], + raw_locations, + raw_value_groups, ) @wrap_with_nvtx_name( @@ -389,28 +473,43 @@ def finish_sparse_delta_refit(self) -> dict[str, Any]: } with torch.no_grad(): - actual = torch.cat( - [ - target.data.view(-1).index_select(0, locations).float() - before - for target, locations, before, _ in verification - ] - ) - expected = torch.cat([item[3] for item in verification]) - difference = (actual - expected).abs() + differences = [] + exact_mismatches = [] + mismatches = [] + samples = 0 + for target, locations, expected_bits in verification: + integer_dtype = sparse_codec.integer_dtype_for_element_size( + target.element_size() + ) + actual_bits = ( + target.data.view(integer_dtype).view(-1).index_select(0, locations) + ) + bit_mismatches = actual_bits.ne(expected_bits) + actual = target.data.view(-1).index_select(0, locations).float() + expected = expected_bits.view(target.dtype).float() + difference = torch.where( + bit_mismatches, (actual - expected).abs(), torch.zeros_like(actual) + ) + differences.append(torch.nan_to_num(difference, nan=float("inf"))) + exact_mismatches.append(bit_mismatches) + mismatches.append( + bit_mismatches + & ~torch.isclose(actual, expected, rtol=1e-6, atol=1e-8) + ) + samples += actual.numel() + difference = torch.cat(differences) stats = torch.stack( ( difference.sum(), difference.max(), - actual.ne(expected).sum().float(), - (~torch.isclose(actual, expected, rtol=1e-6, atol=1e-8)) - .sum() - .float(), + torch.cat(exact_mismatches).sum().float(), + torch.cat(mismatches).sum().float(), ) ).cpu() return { "ok": True, "verification_candidates": candidates, - "verification_samples": actual.numel(), + "verification_samples": samples, "verification_exact_mismatches": int(stats[2]), "verification_mismatches": int(stats[3]), "verification_abs_sum": float(stats[0]), diff --git a/nemo_rl/utils/weight_transfer_sparse_codec.py b/nemo_rl/utils/weight_transfer_sparse_codec.py index aaaf7fa2685..d3b458dec37 100644 --- a/nemo_rl/utils/weight_transfer_sparse_codec.py +++ b/nemo_rl/utils/weight_transfer_sparse_codec.py @@ -17,27 +17,90 @@ import threading from collections.abc import Iterable, Mapping from concurrent.futures import ThreadPoolExecutor -from typing import Any +from typing import Any, Literal import numpy as np import torch NamedTensor = tuple[str, torch.Tensor] TensorBatch = list[NamedTensor] -TensorPayload = tuple[torch.Tensor, torch.Tensor, list[dict[str, Any]]] +SparseOperation = Literal["xor", "overwrite"] +SparseInfo = tuple[str, torch.Tensor, torch.Tensor, torch.Tensor, SparseOperation] +TensorPayload = tuple[torch.Tensor, tuple[torch.Tensor, ...], list[dict[str, Any]]] PreparedTensorPayload = tuple[TensorPayload, int, int] +_INTEGER_DTYPE_BY_SIZE = { + 1: torch.uint8, + 2: torch.int16, + 4: torch.int32, + 8: torch.int64, +} +_DTYPE_BY_NAME = { + "bfloat16": torch.bfloat16, + "float16": torch.float16, + "float32": torch.float32, + "float64": torch.float64, + "float8_e4m3fn": torch.float8_e4m3fn, + "float8_e5m2": torch.float8_e5m2, + "int8": torch.int8, + "int16": torch.int16, + "int32": torch.int32, + "int64": torch.int64, + "uint8": torch.uint8, +} + + +def integer_dtype_for_element_size(element_size: int) -> torch.dtype: + try: + return _INTEGER_DTYPE_BY_SIZE[element_size] + except KeyError as error: + raise ValueError(f"Unsupported tensor element size {element_size}.") from error + + +def dtype_from_name(name: str) -> torch.dtype: + try: + return _DTYPE_BY_NAME[name] + except KeyError as error: + raise ValueError(f"Unsupported sparse-refit tensor dtype {name!r}.") from error + + +def sparse_operation(value: object) -> SparseOperation: + if value == "xor" or value == "overwrite": + return value + raise ValueError(f"Unsupported sparse-refit operation {value!r}.") + + +def _dtype_name(dtype: torch.dtype) -> str: + name = str(dtype).removeprefix("torch.") + if name not in _DTYPE_BY_NAME: + raise ValueError(f"Unsupported sparse-refit tensor dtype {dtype}.") + return name + + +def _integer_view(tensor: torch.Tensor) -> torch.Tensor: + return tensor.contiguous().view( + integer_dtype_for_element_size(tensor.element_size()) + ) + + +def _bytewise_diff_mask(current: torch.Tensor, baseline: torch.Tensor) -> torch.Tensor: + if current.shape != baseline.shape or current.dtype != baseline.dtype: + raise ValueError( + "Current tensor and baseline must have identical shape and dtype." + ) + return _integer_view(current) != _integer_view(baseline) + def encode_sparse_infos( - infos: Iterable[tuple[str, torch.Tensor, torch.Tensor, torch.Tensor]], - *, - empty_dtype: torch.dtype, + infos: Iterable[SparseInfo], ) -> TensorPayload: packed_locations = [] - packed_values = [] + value_parts: list[list[torch.Tensor]] = [] + value_group_by_dtype: dict[torch.dtype, int] = {} + value_offsets: list[int] = [] metadata: list[dict[str, Any]] = [] - index_offset = value_offset = 0 - for name, tensor, raw_locations, raw_values in infos: + index_offset = 0 + for name, tensor, raw_locations, raw_values, operation in infos: count = int(raw_values.numel()) if count == 1 or int(raw_locations[-1] - raw_locations[0] + 1) == count: index_count = 0 @@ -50,27 +113,37 @@ def encode_sparse_infos( location_metadata = {"index_encoding": "deltas"} packed_locations.append(location_tensor) index_count = int(location_tensor.numel()) - packed_values.append(raw_values) + value_group = value_group_by_dtype.get(raw_values.dtype) + if value_group is None: + value_group = len(value_parts) + value_group_by_dtype[raw_values.dtype] = value_group + value_parts.append([]) + value_offsets.append(0) + value_start = value_offsets[value_group] + value_parts[value_group].append(raw_values) + value_offsets[value_group] += count metadata.append( { "name": name, "shape": tuple(int(dim) for dim in tensor.shape), + "dtype": _dtype_name(tensor.dtype), + "operation": operation, "index_start": index_offset, "index_end": index_offset + index_count, - "value_start": value_offset, - "value_end": value_offset + count, + "value_group": value_group, + "value_start": value_start, + "value_end": value_start + count, **location_metadata, } ) index_offset += index_count - value_offset += count indices = ( torch.cat(packed_locations) if packed_locations else torch.empty(0, dtype=torch.uint8) ) - values = ( - torch.cat(packed_values) if packed_values else torch.empty(0, dtype=empty_dtype) + values = tuple( + torch.cat(parts) if len(parts) > 1 else parts[0] for parts in value_parts ) return indices, values, metadata @@ -125,14 +198,7 @@ def __init__(self, config: Mapping[str, Any]) -> None: self.sparse_bucket_size_bytes = int(config["sparse_bucket_size_bytes"]) if self.sparse_bucket_size_bytes < 1: raise ValueError("delta_compression.sparse_bucket_size_bytes must be >= 1") - self.delta_dtype = { - "bf16": torch.bfloat16, - "bfloat16": torch.bfloat16, - "fp16": torch.float16, - "float16": torch.float16, - "fp32": torch.float32, - "float32": torch.float32, - }[str(config["dtype"]).lower()] + self.encoding = sparse_operation(config["encoding"]) self.verification_samples = int( os.getenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "0") ) @@ -160,30 +226,36 @@ def prepare_sparse_delta_payload( baseline = self.baseline.get(name) if baseline is None: raise RuntimeError(f"Sparse delta baseline is missing {name!r}.") - current = tensor.detach().cpu() - current_flat, baseline_flat = current.view(-1), baseline.view(-1) - total_elements += current_flat.numel() - locations = (current_flat != baseline_flat).nonzero().view(-1) + current = tensor.detach().cpu().contiguous() + current_bits = _integer_view(current).view(-1) + baseline_bits = _integer_view(baseline).view(-1) + total_elements += current.numel() + locations = ( + _bytewise_diff_mask(current, baseline).view(-1).nonzero().view(-1) + ) changed_elements += locations.numel() if locations.numel(): - current_values = current_flat[locations] - baseline_values = baseline_flat[locations] - deltas = (current_values - baseline_values).to(self.delta_dtype) - expected_values = baseline_values + deltas.to(baseline.dtype) + current_values = current_bits[locations] + values = ( + current_values.bitwise_xor(baseline_bits[locations]) + if self.encoding == "xor" + else current_values + ) sparse_infos.append( ( name, current, locations, - deltas, + values, + self.encoding, ) ) if self.verification_samples: - verification_sources.append((locations, deltas)) - pending_updates[name] = (locations, expected_values) + verification_sources.append((locations, values)) + pending_updates[name] = (locations, current_values) with self._pending_updates_lock: self._pending_updates.update(pending_updates) - payload = encode_sparse_infos(sparse_infos, empty_dtype=self.delta_dtype) + payload = encode_sparse_infos(sparse_infos) if verification_sources: self._add_verification_samples(payload[2], verification_sources) return payload, changed_elements, total_elements @@ -199,14 +271,14 @@ def _add_verification_samples( ((2 * index + 1) * total) // (2 * count) for index in range(count) ] sample_index = offset = 0 - for item, (locations, deltas) in zip(metadata, sources, strict=True): + for item, (locations, values) in zip(metadata, sources, strict=True): end = offset + locations.numel() while sample_index < count and sample_ranks[sample_index] < end: local_index = sample_ranks[sample_index] - offset location = int(locations[local_index]) item.setdefault("verification_locations", []).append(location) - item.setdefault("verification_deltas", []).append( - float(deltas[local_index]) + item.setdefault("verification_values", []).append( + int(values[local_index]) ) sample_index += 1 offset = end @@ -230,7 +302,10 @@ def on_sync_failed(self) -> None: def snapshot_baseline(self, tensors: Iterable[NamedTensor]) -> None: self._wait_for_baseline_commits() for name, tensor in tensors: - self._baseline(name, tuple(tensor.shape), tensor.dtype).copy_(tensor) + baseline = self._baseline(name, tuple(tensor.shape), tensor.dtype) + baseline.view(torch.uint8).view(-1).copy_( + tensor.detach().cpu().contiguous().view(torch.uint8).view(-1) + ) def _wait_for_baseline_commits(self) -> None: for commit in self._baseline_commits: @@ -238,10 +313,11 @@ def _wait_for_baseline_commits(self) -> None: self._baseline_commits = () def _commit_baseline_updates( - self, updates: Iterable[tuple[str, tuple[torch.Tensor, torch.Tensor]]] + self, + updates: Iterable[tuple[str, tuple[torch.Tensor, torch.Tensor]]], ) -> None: for name, (locations, values) in updates: - target = self.baseline[name].view(-1) + target = _integer_view(self.baseline[name]).view(-1) count = locations.numel() if count > 1: first, last = int(locations[0]), int(locations[-1]) @@ -269,16 +345,18 @@ def _baseline( ) -> torch.Tensor: if name in self.baseline: return self.baseline[name] + numel = torch.Size(shape).numel() + nbytes = numel * torch.empty((), dtype=dtype).element_size() if self.baseline_in_memory: - baseline = torch.empty(shape, dtype=dtype) + storage = torch.empty(nbytes, dtype=torch.uint8) else: - numel = torch.Size(shape).numel() with tempfile.NamedTemporaryFile( prefix="nrl-refit-baseline-", dir=self.baseline_mmap_dir ) as handle: - handle.truncate(numel * torch.empty((), dtype=dtype).element_size()) - baseline = torch.from_file( - handle.name, shared=True, size=numel, dtype=dtype - ).view(shape) + handle.truncate(nbytes) + storage = torch.from_file( + handle.name, shared=True, size=nbytes, dtype=torch.uint8 + ) + baseline = storage.view(dtype).view(shape) self.baseline[name] = baseline return baseline diff --git a/tests/unit/models/generation/test_vllm_sparse_delta.py b/tests/unit/models/generation/test_vllm_sparse_delta.py index 49b2c2bcbcb..757f16f9082 100644 --- a/tests/unit/models/generation/test_vllm_sparse_delta.py +++ b/tests/unit/models/generation/test_vllm_sparse_delta.py @@ -13,7 +13,6 @@ # limitations under the License. import math -import sys from types import MethodType, SimpleNamespace from typing import Any @@ -24,7 +23,10 @@ VllmSparseDeltaApplier, _SparseLoadTracer, ) -from nemo_rl.utils.weight_transfer_sparse_codec import encode_sparse_infos +from nemo_rl.utils.weight_transfer_sparse_codec import ( + encode_sparse_infos, + integer_dtype_for_element_size, +) class _NativeLoaderModel: @@ -42,6 +44,8 @@ def load_weights(self, weights): target = self.targets if name == "weight": target["identity"].copy_(source) + elif name == "weight_scale_inv": + target["scale"].copy_(source) elif name.endswith("self_attn.k_proj.weight"): target["qkv"][4:6].copy_(source[2:4]) elif name.endswith("mlp.gate_proj.weight"): @@ -73,11 +77,9 @@ def _applier(model: Any) -> VllmSparseDeltaApplier: ) -def _stub_fp8(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setitem( - sys.modules, - "nemo_rl.models.generation.vllm.quantization.fp8", - SimpleNamespace(is_fp8_model=lambda _config: False), +def _bits(values: torch.Tensor) -> torch.Tensor: + return values.contiguous().view( + integer_dtype_for_element_size(values.element_size()) ) @@ -85,7 +87,11 @@ def _stub_fp8(monkeypatch: pytest.MonkeyPatch) -> None: def test_serialized_sparse_payload_batch_preserves_order(tmp_path) -> None: applier = VllmSparseDeltaApplier(SimpleNamespace(), torch.device("cpu")) payloads = [ - (torch.tensor([index]), torch.tensor([float(index)]), [{"index": index}]) + ( + torch.tensor([index]), + (torch.tensor([float(index)]),), + [{"index": index}], + ) for index in range(3) ] paths = [tmp_path / f"{index}.pt" for index in range(3)] @@ -113,8 +119,7 @@ def test_serialized_sparse_payload_batch_preserves_order(tmp_path) -> None: @pytest.mark.vllm -def test_native_loaders_compile_sparse_placement(monkeypatch) -> None: - _stub_fp8(monkeypatch) +def test_native_loaders_compile_sparse_placement() -> None: targets = { "identity": torch.zeros(4), "qkv": torch.zeros(8, 2), @@ -141,7 +146,7 @@ def test_native_loaders_compile_sparse_placement(monkeypatch) -> None: "backbone.layers.0.mixer.A_log", (2,), [0, 1], - [math.log(1.5), math.log(0.5)], + [math.log(3.0), math.log(2.0)], ), ("model.layers.0.mlp.experts.7.gate_proj.weight", (8, 2), [8], [9]), ] @@ -151,16 +156,26 @@ def test_native_loaders_compile_sparse_placement(monkeypatch) -> None: name, torch.empty(shape), torch.tensor(locations), - torch.tensor(values, dtype=torch.float32), + _bits(torch.tensor(values, dtype=torch.float32)), + "overwrite", ) for name, shape, locations, values in infos + ] + ) + payload[2][-1].update( + verification_locations=[8], + verification_values=[int(_bits(torch.tensor([9.0]))[0])], + ) + payload[2][7].update( + verification_locations=[0, 1], + verification_values=[ + int(value) for value in _bits(torch.tensor([math.log(3.0), math.log(2.0)])) ], - empty_dtype=torch.float32, ) - payload[2][-1].update(verification_locations=[8], verification_deltas=[9.0]) applier = _applier(_NativeLoaderModel(**targets)) applier._apply_sparse_weight_deltas(payload[:2], payload[2]) + verification = applier.finish_sparse_delta_refit() assert torch.equal(targets["identity"], torch.tensor([0.0, 1.0, 2.0, 0.0])) assert targets["qkv"].view(-1)[[8, 9, 11]].tolist() == [1.0, 1.0, 1.0] @@ -169,7 +184,74 @@ def test_native_loaders_compile_sparse_placement(monkeypatch) -> None: assert targets["w2"].view(-1)[[8, 11, 12, 15]].tolist() == [5, 5, 5, 5] assert targets["mamba"].view(-1)[[0, 3, 4, 11]].tolist() == [6, 6, 6, 6] assert torch.allclose(targets["a"], torch.tensor([-3.0, -2.0])) - assert (applier._verification_candidates, applier._verification) == (1, []) + assert verification["verification_candidates"] == 3 + assert verification["verification_samples"] == 2 + assert verification["verification_exact_mismatches"] == 0 + + +@pytest.mark.vllm +def test_xor_applies_through_packed_native_loaders() -> None: + targets = { + "qkv": torch.zeros(8, 2), + "merged": torch.zeros(8, 2), + "w13": torch.zeros(2, 4, 2), + "mamba": torch.zeros(6, 2), + } + infos = [ + ( + "model.layers.0.self_attn.k_proj.weight", + (4, 2), + [0, 4, 5, 7], + [1.0] * 4, + ), + ( + "model.layers.0.mlp.gate_proj.weight", + (8, 2), + [0, 8, 9, 15], + [2.0] * 4, + ), + ( + "model.layers.0.mlp.up_proj.weight", + (8, 2), + [8, 15], + [3.0] * 2, + ), + ( + "model.layers.0.mlp.experts.3.gate_proj.weight", + (8, 2), + [8, 15], + [4.0] * 2, + ), + ( + "model.layers.0.mixer.in_proj.weight", + (10, 2), + [0, 4, 7, 12, 19], + [6.0] * 5, + ), + ] + payload = encode_sparse_infos( + [ + ( + name, + torch.empty(shape), + torch.tensor(locations), + _bits(torch.tensor(values)).bitwise_xor( + _bits(torch.zeros(len(values))) + ), + "xor", + ) + for name, shape, locations, values in infos + ] + ) + + _applier(_NativeLoaderModel(**targets))._apply_sparse_weight_deltas( + payload[:2], payload[2] + ) + + assert targets["qkv"].view(-1)[[8, 9, 11]].tolist() == [1.0, 1.0, 1.0] + assert targets["merged"].view(-1)[[0, 1, 7, 8, 15]].tolist() == [2, 2, 2, 3, 3] + assert targets["w13"].view(-1)[[8, 15]].tolist() == [4.0, 4.0] + assert targets["mamba"].view(-1)[[0, 3, 4, 11]].tolist() == [6, 6, 6, 6] @pytest.mark.vllm @@ -264,39 +346,78 @@ def trace(target, source_shape, load): @pytest.mark.vllm -def test_unknown_native_loader_fails_closed(monkeypatch) -> None: - _stub_fp8(monkeypatch) +def test_unknown_native_loader_fails_closed() -> None: model = _NativeLoaderModel(identity=torch.zeros(1)) for name, error in (("unknown", "did not place"), ("transformed", "transform")): payload = encode_sparse_infos( - [(name, torch.empty(1), torch.tensor([0]), torch.tensor([1.0]))], - empty_dtype=torch.float32, + [ + ( + name, + torch.empty(1), + torch.tensor([0]), + _bits(torch.tensor([1.0])), + "overwrite", + ) + ], ) with pytest.raises(RuntimeError, match=error): _applier(model)._apply_sparse_weight_deltas(payload[:2], payload[2]) +@pytest.mark.vllm +def test_unknown_sparse_operation_fails_closed() -> None: + target = torch.zeros(1) + payload = encode_sparse_infos( + [ + ( + "weight", + target, + torch.tensor([0]), + _bits(torch.tensor([1.0])), + "overwrite", + ) + ] + ) + payload[2][0]["operation"] = "unknown" + + with pytest.raises(ValueError, match="Unsupported sparse-refit operation"): + _applier(_NativeLoaderModel(identity=target))._apply_sparse_weight_deltas( + payload[:2], payload[2] + ) + + @pytest.mark.vllm @pytest.mark.parametrize( - ("initial", "expected_delta", "exact_mismatches", "mismatches"), + ("initial", "verified_value", "exact_mismatches", "mismatches"), [(200.0, 4.0, 0, 0), (2.0, 4.0000005, 1, 0), (2.0, 5.0, 1, 1)], ) -def test_sparse_delta_verification_compares_applied_delta( - monkeypatch, +def test_sparse_delta_verification_compares_replacement( initial: float, - expected_delta: float, + verified_value: float, exact_mismatches: int, mismatches: int, ) -> None: - _stub_fp8(monkeypatch) target = torch.tensor([1.0, initial, 3.0, initial]) + replacement = torch.tensor([initial + 4.0, initial + 4.0]) payload = encode_sparse_infos( - [("weight", target, torch.tensor([1, 3]), torch.tensor([4.0, 4.0]))], - empty_dtype=target.dtype, + [ + ( + "weight", + target, + torch.tensor([1, 3]), + _bits(replacement), + "overwrite", + ) + ], ) payload[2][0].update( verification_locations=[1, 3], - verification_deltas=[expected_delta, expected_delta], + verification_values=[ + int(value) + for value in _bits( + torch.tensor([initial + verified_value, initial + verified_value]) + ) + ], ) applier = _applier(_NativeLoaderModel(identity=target)) @@ -308,3 +429,166 @@ def test_sparse_delta_verification_compares_applied_delta( assert result["verification_samples"] == 2 assert result["verification_exact_mismatches"] == 2 * exact_mismatches assert result["verification_mismatches"] == 2 * mismatches + + +@pytest.mark.vllm +def test_fp8_weight_and_scale_use_exact_bit_overwrite() -> None: + target = torch.tensor([0x38, 0x40, 0x48], dtype=torch.uint8).view( + torch.float8_e4m3fn + ) + scale = torch.tensor([1.0, 2.0]) + current = target.clone() + current.view(torch.uint8)[1] = 0x41 + current.view(torch.uint8)[2] = 0x7F + current_scale = scale.clone() + current_scale[0] = 1.5 + payload = encode_sparse_infos( + [ + ( + "weight", + current, + torch.tensor([1, 2]), + current.view(torch.uint8)[1:3], + "overwrite", + ), + ( + "weight_scale_inv", + current_scale, + torch.tensor([0]), + current_scale.view(torch.int32)[:1], + "overwrite", + ), + ] + ) + payload[2][0].update( + verification_locations=[1, 2], verification_values=[0x41, 0x7F] + ) + payload[2][1].update( + verification_locations=[0], + verification_values=[int(current_scale.view(torch.int32)[0])], + ) + applier = _applier(_NativeLoaderModel(identity=target, scale=scale)) + + applier._apply_sparse_weight_deltas(payload[:2], payload[2]) + applier._apply_sparse_weight_deltas(payload[:2], payload[2]) + result = applier.finish_sparse_delta_refit() + + assert target.view(torch.uint8).tolist() == [0x38, 0x41, 0x7F] + assert torch.equal(scale, current_scale) + assert result["verification_samples"] == 6 + assert result["verification_exact_mismatches"] == 0 + assert result["verification_mismatches"] == 0 + assert result["verification_abs_sum"] == 0.0 + + +@pytest.mark.vllm +def test_xor_applies_exact_bits_and_replay_reverts() -> None: + baseline = torch.tensor([1.0, 2.0, 3.0]) + target = baseline.clone() + current = torch.tensor([1.0, 5.0, -0.0]) + locations = torch.tensor([1, 2]) + xor_values = _bits(current)[locations].bitwise_xor(_bits(baseline)[locations]) + payload = encode_sparse_infos([("weight", current, locations, xor_values, "xor")]) + payload[2][0].update( + verification_locations=locations.tolist(), + verification_values=[int(value) for value in xor_values], + ) + applier = _applier(_NativeLoaderModel(identity=target)) + + applier._apply_sparse_weight_deltas(payload[:2], payload[2]) + result = applier.finish_sparse_delta_refit() + + assert torch.equal(_bits(target), _bits(current)) + assert result["verification_exact_mismatches"] == 0 + assert result["verification_mismatches"] == 0 + + applier._apply_sparse_weight_deltas(payload[:2], payload[2]) + assert torch.equal(_bits(target), _bits(baseline)) + + +@pytest.mark.vllm +def test_overwrite_casts_absolute_source_values() -> None: + target = torch.zeros(2, dtype=torch.float16) + source = torch.tensor([1.25, -2.5], dtype=torch.float32) + payload = encode_sparse_infos( + [ + ( + "weight", + source, + torch.tensor([0, 1]), + _bits(source), + "overwrite", + ) + ] + ) + + _applier(_NativeLoaderModel(identity=target))._apply_sparse_weight_deltas( + payload[:2], payload[2] + ) + + assert torch.equal(target, source.to(torch.float16)) + + +@pytest.mark.vllm +@pytest.mark.parametrize( + ("name", "source", "targets", "error"), + [ + ( + "backbone.layers.0.mixer.A_log", + torch.tensor([math.log(2.0)]), + {"a": torch.tensor([-1.0])}, + "transformed weight", + ), + ( + "weight", + torch.tensor([1.0], dtype=torch.float32), + {"identity": torch.zeros(1, dtype=torch.float16)}, + "dtypes differ", + ), + ], +) +def test_xor_rejects_non_bitwise_compatible_targets( + name: str, + source: torch.Tensor, + targets: dict[str, torch.Tensor], + error: str, +) -> None: + payload = encode_sparse_infos( + [(name, source, torch.tensor([0]), _bits(source), "xor")] + ) + + with pytest.raises(RuntimeError, match=error): + _applier(_NativeLoaderModel(**targets))._apply_sparse_weight_deltas( + payload[:2], payload[2] + ) + + +@pytest.mark.vllm +def test_xor_rejects_overlapping_target_mappings() -> None: + class RepeatedCopyModel(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.weight = torch.nn.Parameter(torch.zeros(2)) + + def load_weights(self, weights) -> None: + for _, source in weights: + self.weight.copy_(source) + self.weight.copy_(source) + + source = torch.tensor([1.0, 2.0]) + payload = encode_sparse_infos( + [ + ( + "weight", + source, + torch.tensor([0, 1]), + _bits(source).bitwise_xor(_bits(torch.zeros_like(source))), + "xor", + ) + ] + ) + + with pytest.raises(RuntimeError, match="target mappings overlap"): + _applier(RepeatedCopyModel())._apply_sparse_weight_deltas( + payload[:2], payload[2] + ) diff --git a/tests/unit/utils/test_weight_transfer_remote_sparse.py b/tests/unit/utils/test_weight_transfer_remote_sparse.py index 3add0c4ff31..cac740aa5a2 100644 --- a/tests/unit/utils/test_weight_transfer_remote_sparse.py +++ b/tests/unit/utils/test_weight_transfer_remote_sparse.py @@ -25,6 +25,7 @@ from nemo_rl.utils.weight_transfer_remote_sparse import download_s3_refit_payload from nemo_rl.utils.weight_transfer_sparse_codec import ( DeltaCompressionTracker, + _bytewise_diff_mask, encode_sparse_infos, sparse_locations_for_item, ) @@ -59,8 +60,10 @@ def _stream_sparse_test_payloads(tensors, send_payload): ) -def _delta_tracker() -> DeltaCompressionTracker: - return DeltaCompressionTracker({"dtype": "bf16", "sparse_bucket_size_bytes": 1024}) +def _delta_tracker(encoding: str = "overwrite") -> DeltaCompressionTracker: + return DeltaCompressionTracker( + {"encoding": encoding, "sparse_bucket_size_bytes": 1024} + ) def test_delta_tracker_commits_only_successful_syncs(monkeypatch) -> None: @@ -90,37 +93,138 @@ def test_delta_tracker_emits_bounded_delta_samples(monkeypatch) -> None: ) assert metadata[0]["verification_locations"] == [1, 3] - assert metadata[0]["verification_deltas"] == [1.0, 1.0] + assert metadata[0]["verification_values"] == [ + int(tensor.view(torch.int32)[1]), + int(tensor.view(torch.int32)[3]), + ] + assert metadata[0]["operation"] == "overwrite" assert (changed, total) == (2, 4) -def test_delta_tracker_commits_quantized_receiver_baseline(monkeypatch) -> None: +def test_delta_tracker_commits_exact_source_baseline(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") tracker = _delta_tracker() tensor = torch.tensor([1.0]) tracker.snapshot_baseline([("weight", tensor)]) tensor.add_(0.001) - (_, deltas, _), _, _ = tracker.prepare_sparse_delta_payload([("weight", tensor)]) - expected = torch.tensor([1.0]) + deltas.float() + (_, value_groups, _), _, _ = tracker.prepare_sparse_delta_payload( + [("weight", tensor)] + ) + assert torch.equal(value_groups[0], tensor.view(torch.int32)) tracker.on_sync_succeeded() tracker.prepare_sparse_delta_payload([("weight", tensor)]) - assert torch.equal(tracker.baseline["weight"], expected) - assert not torch.equal(expected, tensor) + assert torch.equal(tracker.baseline["weight"], tensor) + + +def test_delta_tracker_xor_encodes_against_baseline(monkeypatch) -> None: + monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") + tracker = _delta_tracker("xor") + tensor = torch.tensor([1.0, 2.0, 3.0]) + baseline = tensor.clone() + tracker.snapshot_baseline([("weight", tensor)]) + tensor[[0, 2]] = torch.tensor([4.0, 5.0]) + + (_, value_groups, metadata), changed, total = tracker.prepare_sparse_delta_payload( + [("weight", tensor)] + ) + locations = torch.tensor([0, 2]) + expected = tensor.view(torch.int32)[locations].bitwise_xor( + baseline.view(torch.int32)[locations] + ) + + assert torch.equal(value_groups[0], expected) + assert metadata[0]["operation"] == "xor" + assert (changed, total) == (2, 3) + tracker.on_sync_succeeded() + tracker.prepare_sparse_delta_payload([("weight", tensor)]) + assert torch.equal(tracker.baseline["weight"], tensor) def test_sparse_index_encoding_preserves_uint64_locations() -> None: locations = torch.tensor([0, 2**32 + 5]) packed, _, metadata = encode_sparse_infos( - [("weight", torch.empty(2), locations, torch.ones(2))], - empty_dtype=torch.float32, + [ + ( + "weight", + torch.empty(2), + locations, + torch.ones(2, dtype=torch.int32), + "overwrite", + ) + ], ) decoded = sparse_locations_for_item(metadata[0], packed, device="cpu") assert torch.equal(decoded, locations) +def test_bytewise_diff_mask_supports_float8() -> None: + baseline = torch.tensor([0x38, 0x7F, 0x00], dtype=torch.uint8).view( + torch.float8_e4m3fn + ) + current = baseline.clone() + current.view(torch.uint8)[1] = 0x7E + + assert _bytewise_diff_mask(current, baseline).tolist() == [False, True, False] + + +@pytest.mark.parametrize("encoding", ["xor", "overwrite"]) +def test_delta_tracker_encodes_fp8_weight_and_scale_bits( + monkeypatch, encoding: str +) -> None: + monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") + monkeypatch.setenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "2") + tracker = _delta_tracker(encoding) + weight = torch.tensor([0x38, 0x40, 0x48], dtype=torch.uint8).view( + torch.float8_e4m3fn + ) + scale = torch.tensor([1.0, 2.0], dtype=torch.float32) + baseline_weight = weight.clone() + baseline_scale = scale.clone() + tracker.snapshot_baseline([("weight", weight), ("weight_scale_inv", scale)]) + weight.view(torch.uint8)[1] = 0x41 + scale[0] = 1.5 + + (locations, value_groups, metadata), changed, total = ( + tracker.prepare_sparse_delta_payload( + [("weight", weight), ("weight_scale_inv", scale)] + ) + ) + + assert (changed, total) == (2, 5) + assert len(value_groups) == 2 + assert [item["operation"] for item in metadata] == [encoding, encoding] + assert [item["dtype"] for item in metadata] == ["float8_e4m3fn", "float32"] + expected_weight = int(weight.view(torch.uint8)[1]) + expected_scale = int(scale.view(torch.int32)[0]) + if encoding == "xor": + expected_weight ^= int(baseline_weight.view(torch.uint8)[1]) + expected_scale ^= int(baseline_scale.view(torch.int32)[0]) + assert [item["verification_values"] for item in metadata] == [ + [expected_weight], + [expected_scale], + ] + assert [ + sparse_locations_for_item(item, locations, device="cpu").tolist() + for item in metadata + ] == [ + [1], + [0], + ] + + tracker.on_sync_succeeded() + assert not tracker.prepare_sparse_delta_payload( + [("weight", weight), ("weight_scale_inv", scale)] + )[0][2] + + +def test_delta_tracker_rejects_arithmetic_encoding() -> None: + with pytest.raises(ValueError, match="Unsupported sparse-refit operation"): + _delta_tracker("add") + + def test_s3_download_verifies_checksum(monkeypatch) -> None: compressed = zstandard.ZstdCompressor().compress(b"payload") monkeypatch.setattr( From 75b1dc9c77fcc1beada1b44cdfd41835be644284 Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Mon, 13 Jul 2026 00:35:36 -0700 Subject: [PATCH 08/18] Use mcore local baseline Signed-off-by: Hollow Man --- docs/design-docs/sparse-delta-refit.md | 232 +++++++-- nemo_rl/algorithms/grpo.py | 42 +- .../models/generation/vllm/vllm_backend.py | 22 +- .../generation/vllm/vllm_sparse_delta.py | 242 +++++---- .../generation/vllm/vllm_sparse_refit.py | 253 ++++++++-- nemo_rl/models/generation/vllm/vllm_worker.py | 12 +- .../generation/vllm/vllm_worker_async.py | 5 - .../policy/workers/megatron_policy_worker.py | 44 +- .../workers/megatron_remote_sparse_refit.py | 467 +++++++++++++++++- .../utils/weight_transfer_remote_sparse.py | 219 ++++++-- nemo_rl/utils/weight_transfer_sparse_codec.py | 457 +++++++++++++---- nemo_rl/utils/weight_transfer_zmq.py | 61 +-- .../vllm_remote_sparse_weight_synchronizer.py | 52 +- pyrefly.toml | 14 +- .../generation/test_vllm_sparse_delta.py | 273 ++++++++-- .../generation/test_vllm_sparse_refit.py | 185 +++++-- .../test_megatron_remote_sparse_refit.py | 432 +++++++++++++++- .../test_weight_transfer_remote_sparse.py | 209 +++++++- ..._vllm_remote_sparse_weight_synchronizer.py | 99 +++- tools/refit_bandwidth_calculator.py | 6 +- 20 files changed, 2759 insertions(+), 567 deletions(-) diff --git a/docs/design-docs/sparse-delta-refit.md b/docs/design-docs/sparse-delta-refit.md index d0022be6e9e..0638cc2c413 100644 --- a/docs/design-docs/sparse-delta-refit.md +++ b/docs/design-docs/sparse-delta-refit.md @@ -1,10 +1,12 @@ # Remote Sparse-Delta vLLM Refit Remote sparse-delta refit updates non-colocated vLLM workers without sending a -full checkpoint after every optimizer step. Megatron workers export Hugging -Face (HF) weights, compare them with a sharded CPU baseline, and send changed -locations and byte-encoded values through S3 or ZeroMQ. vLLM maps those HF -coordinates into its local TP and EP layouts and applies the updates in place. +full checkpoint after every optimizer step. Megatron workers compare every +uniquely owned MCore tensor against a policy-local CPU baseline. Exact affine +mappings emit sparse Hugging Face (HF) coordinates directly; only changed tasks +whose conversion is not affine traverse Megatron Bridge. S3 or ZeroMQ carries +the resulting payloads, and vLLM maps the HF coordinates into its local TP and +EP layout before applying them in place. The feature is opt-in. Its synchronizer, codec, transports, receiver queue, and placement engine are separate from existing NCCL, CUDA IPC, and packed refit @@ -31,11 +33,15 @@ the global flush completes. ```mermaid flowchart LR subgraph P["Megatron policy cluster"] - B["Megatron Bridge HF export"] - C["Sharded CPU or mmap baseline"] + T["Bridge conversion-task metadata"] + L["All unique MCore shards"] + B["Changed residual Bridge export"] + C["Local source and sharded HF baselines"] E["Compare, encode, and compress"] + T --> L + T --> B + L --> C B --> C - B --> E C --> E end @@ -44,7 +50,7 @@ flowchart LR subgraph G["vLLM generation cluster"] H["HTTP receiver"] - Q["Bounded FIFO apply queue"] + Q["Eager node staging and bounded FIFO apply queue"] A["Sparse placement and apply"] H --> Q --> A end @@ -70,12 +76,39 @@ and commit protocol.* ### Initialize the baseline -`VllmRemoteSparseWeightSynchronizer.init_communicator()` starts baseline -construction on every policy worker and then discovers the vLLM HTTP endpoints. -Each worker participates in the HF export but stores only chunks assigned by -`chunk_index % shard_count`. The baseline is therefore sharded across policy -workers. It uses file-backed `torch.from_file` tensors by default; -`NRL_REFIT_BASELINE_IN_MEMORY=1` keeps it in RAM. +The policy initialization task starts baseline construction as soon as its +workers are ready, while the independent vLLM model load continues. +`VllmRemoteSparseWeightSynchronizer.init_communicator()` discovers the receiver +endpoints and then joins the prelaunched baseline before setup returns. The +first rollout therefore does not enter a redundant weight sync or race an +unfinished snapshot with policy training. Conversion tasks are split into two +deterministic paths without changing Megatron Bridge: + +- Every conversion task keeps its source baseline in MCore layout. A stable name + hash assigns replicated dense tensors across their combined DP/CP and TP + replicas, and expert tensors across their expert-DP replicas. TP, PP, EP, and + ETP still contribute their unique shards. This keeps exactly one source copy + while sharing baseline scans and uploads across equivalent ranks. +- Exact direct, column, row, replicated, and gated mappings use that baseline + to produce HF-coordinate deltas without Bridge export. The decision is based + on the resolved Bridge mapping type, not a model name or weight suffix. +- Tasks with a transform also keep a canonical HF baseline, sharded across + workers by a stable hash of the HF name. Stable ownership is required because + later refits export only changed residual tasks and therefore have a different + chunk sequence. + +The local baselines contain one copy of each unique source element across the +policy workers, rather than a full HF copy per DP replica. Residual tasks have +one additional canonical HF copy distributed across the workers. Baselines use +file-backed `torch.from_file` tensors by default; +`NRL_REFIT_BASELINE_IN_MEMORY=1` keeps them in RAM. Local snapshotting and the +residual Bridge export run concurrently. + +Baseline initialization also returns each canonical tensor's name, shape, and +dtype. The synchronizer merges that metadata and asks every vLLM worker to +compile its native weight-loader placement plan before the first transfer. +This does not export HF values or mutate vLLM weights; it moves loader tracing +and plan validation out of the first timed refit. On a fresh run, vLLM already holds the shared checkpoint. Baseline construction starts early and can overlap initial generation, so the redundant initial full @@ -83,24 +116,67 @@ sync is skipped. On resume, both clusters must still start from the same HF weight version; sparse refit does not reconstruct a rollout baseline from an arbitrary training checkpoint. -### Export and encode deltas +### Compare and encode deltas + +Every uniquely owned local tensor is copied to CPU and compared bytewise. For +an affine mapping, changed flat locations are mapped into the unsharded HF +tensor using the TP or ETP rank, shard dimension, and EP-global expert number. +This covers more than FFNs: attention output projections, Mamba affine weights, +norms, routers, shared experts, and other exact mappings use the same path. + +For non-affine tasks, one integer flag per conversion task is reduced across the +policy world. Bridge exports only globally changed tasks, after which the +canonical HF tracker computes the sparse payload. If one member of a grouped +export changes, the complete group is exported. A model-specific Bridge +postprocessor can have undeclared cross-task dependencies, so such bridges keep +the full residual task set whenever any residual task changes. Compound QKV, +Mamba packing, permutations, grouped or fused exports, padded or tied +embeddings, and other custom transformations therefore retain Bridge semantics. + +Policy-local comparison removes full-tensor TP/EP gathers and PP broadcasts for +directly projectable weights, but comparison itself is still proportional to +model size: every unique local element is copied to CPU and scanned. Let `P` be +directly projectable bytes, `R` residual source bytes, `R_changed` the residual +HF tensors selected by the task flags, and `s` the element change fraction. The +leading work is approximately: -Every refit still traverses `MegatronBridge.export_hf_weights()`. Sharding the -baseline avoids duplicate baseline storage and payload production, but it does -not remove Bridge export or the full CPU comparison. Low changed density mainly -reduces encoded and transferred bytes. +```text +old: Bridge(P + R) + HF D2H/scan(P + R) + wire(s(P + R)) +new: local D2H/scan(P + R) + Bridge(R_changed) + HF scan(R_changed) + + wire(s(P + R)) + one O(number_of_tasks) flag all-reduce +``` + +When Adam changes at least one element in every residual tensor, +`R_changed` approaches `R`; the gain then comes from removing Bridge work for +`P`, not from task filtering. It helps less when local D2H or CPU scanning is +the bottleneck, most bytes use custom transformations, or `s` is high. The +reported changed percentage is computed from unique policy-local source +elements, so the extra detector does not inflate it with a second HF scan. For each assigned chunk, `DeltaCompressionTracker` finds changed flat locations through an integer view with the same element width. It encodes either the absolute new bits (`overwrite`) or new bits XOR baseline bits -(`xor`). Both encodings are dtype-blind and preserve FP8 bit patterns. +(`xor`). Both encodings are dtype-blind and preserve FP8 bit patterns in +the codec, although end-to-end FP8 rollout refit is outside the supported scope. `overwrite` is idempotent and recommended. `xor` can compress better, but it requires an exact same-dtype receiver baseline and exactly-once application. -The pending source baseline always records the exact new source bits. - -The producer overlaps export, encoding, `torch.save` serialization, zstd level -1 compression, and transfer with fixed-size executors. Source baselines do not -commit until the entire transfer succeeds. +Selecting `xor` enables mixed operation: directly projected, bitwise-compatible +policy shards use XOR, while Bridge residuals and the full-HF compatibility +path use overwrite. A payload batch may therefore contain both operations. The +receiver validates direct-loader compatibility and fails closed on a transform, +dtype cast, or overlapping mapping; use `overwrite` for models with any such +direct loader. The pending source baseline always records the exact new source +bits. + +The producer pulls bounded export chunks and compares them in parallel. A +separate bounded stage coalesces encoded chunks up to +`sparse_bucket_size_bytes`, serializes them, and applies zstd level 1 before the +transport executor. Separating the 256 MiB S3 compare chunk from the 1 GiB wire +bucket preserves D2H/scan parallelism while reducing object and manifest count. +The stages run concurrently, so payload N transfers while later chunks are +compared and encoded. Source baselines do not commit until the entire transfer +succeeds. Worker errors are reported only after every rank drains its Bridge +export iterator; stopping early can strand peers in a conversion collective. ### Transfer and apply @@ -116,10 +192,31 @@ throughput target. ZeroMQ assigns each producer to one relay; that relay fans the compressed payload out to every generation replica. Both transports use the same receiver endpoints and checksum validation. -The receiver deduplicates payload identities, batches them in a bounded FIFO -queue, and applies batches on one worker thread. When all vLLM ranks share a -node, payloads are staged under `/dev/shm` and passed to collective RPC by file -path. Otherwise, the serialized batch is sent through one collective RPC. +The receiver deduplicates payload identities and applies bounded batches on one +FIFO worker thread. Each generation replica downloads a transport payload +once. When its vLLM ranks share a node, decode and flat-file staging under +`/dev/shm` begin as soon as each payload arrives, without waiting for the batch +to fill. Staging futures then feed the serial collective apply worker, so +download, decompression, staging, and earlier GPU applies can overlap. Queue +depth provides backpressure; the default depth and batch size bound the pending +window at 256 payloads. + +Locations use `int32` unless a single canonical tensor exceeds the signed +32-bit index range; values remain grouped by dtype. The collective RPC passes +only file paths, and workers use `torch.load(..., mmap=True)`. The mmap is not +a second baseline: it lets eight colocated ranks share the staged file's page +cache instead of materializing eight independent CPU copies, while each rank +copies only its selected entries to its GPU. + +Each worker selects the canonical entries consumed by its TP/EP placement plan +on CPU, converts only those selected locations to CUDA `int64`, and then copies +the selected locations and values to its GPU. Thus transport bytes are not +duplicated within a node and irrelevant canonical entries do not cross H2D. +The staged format flattens locations into one `int32` and one `int64` tensor; +it does not serialize one tensor object per model parameter. If ranks do not +share a node, the receiver still decodes once and sends that flat +representation through one collective RPC. Every worker derives and caches its +source plan from the native loader before CPU partitioning. The final `/nemo-rl/refit/flush` drains the queue, synchronizes CUDA, and checks optional delta samples. Only then does the source commit pending baseline @@ -153,12 +250,20 @@ later sparse values directly with bitwise XOR or `index_copy_`, without materializing a dense HF tensor or duplicating QKV, MoE, Mamba, or TP placement rules. +The same recorded copies define a source plan for node-local partitioning. A +source plan contains the canonical offset, shape, and strides consumed by one +worker. Linear routes use sorted-range lookup; strided routes use affine source +coordinates. Verification samples are partitioned by the same plan. An unknown +or transformed loader fails while compiling the plan, before the receiver +stages or applies that tensor. + The tracer accepts affine tensor views. Absolute overwrite also supports the -Mamba `A_log` transform and source-to-target dtype casts. XOR rejects either -case because target bits no longer match the source representation. XOR also -rejects overlapping target mappings. An element-expanding copy, unknown -transform, or unplaced non-expert tensor fails before any payload update. There -is no dense fallback for an unknown layout. +Mamba `A_log` transform and source-to-target dtype casts. Bridge residuals use +overwrite even when `encoding: xor` is configured because their target bits may +not match the HF source representation. XOR also rejects overlapping target +mappings. An element-expanding copy, unknown transform, or unplaced non-expert +tensor fails before any payload update. There is no dense fallback for an +unknown layout. ## Configuration @@ -170,8 +275,8 @@ policy: backend: vllm refit_transport: vllm_s3_sparse # or vllm_zmq_sparse delta_compression: - encoding: overwrite # xor requires an exact baseline and exactly-once apply - sparse_bucket_size_bytes: 268435456 + encoding: overwrite # xor requires bitwise-compatible direct vLLM loaders + sparse_bucket_size_bytes: 1073741824 colocated: enabled: false vllm_cfg: @@ -192,19 +297,21 @@ must contain the same nonempty token on producers and receivers. | Control | Default | |---|---:| | `NRL_REFIT_S3_EXPORT_CHUNK_BYTES` | 256 MiB | -| `NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES` | 1 GiB | +| `NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES` | 256 MiB | | `NRL_REFIT_{S3,ZMQ}_ENCODE_WORKERS` | 2-8 from CPU count | | `NRL_REFIT_S3_UPLOAD_WORKERS` | 4-32 from CPU count | | `NRL_REFIT_ZMQ_SEND_WORKERS` | 4 | | `NRL_REFIT_ZMQ_RELAY_PAYLOAD_WORKERS` | 16 | | `NRL_REFIT_ZMQ_RELAY_FANOUT_WORKERS` | 8-32 from replica count | -| `NRL_REFIT_APPLY_QUEUE_DEPTH` / `NRL_REFIT_APPLY_BATCH_SIZE` | 2 / 8 | +| `NRL_REFIT_APPLY_QUEUE_DEPTH` / `NRL_REFIT_APPLY_BATCH_SIZE` | 32 / 8 | +| `NRL_REFIT_PARTITION_WORKERS` | 2-8 from CPU count | | `NRL_REFIT_{S3,ZMQ}_ZSTD_THREADS` | 0 | | `NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD` | 0 | -Export chunks are also capped by `sparse_bucket_size_bytes` and the packed -tensor limit. Increase one concurrency control at a time; excessive parallelism -can move the bottleneck into host memory, collective export, relay fanout, or +Export chunks are capped by `sparse_bucket_size_bytes` and the packed tensor +limit, but they intentionally remain smaller than the recommended S3 wire +bucket. Increase one concurrency control at a time; excessive parallelism can +move the bottleneck into host memory, collective export, relay fanout, or receiver apply. ## Metrics and profiling @@ -212,15 +319,26 @@ receiver apply. | Signal | Meaning | |---|---| | `REFIT_BASELINE_INIT` | Baseline export and snapshot time | +| `REFIT_RECEIVER_PREWARM` | Native vLLM placement plans compiled during initialization | | `REFIT_{S3,ZMQ}_TIMING` | Producer wall time, stage service time, payloads, bytes, and changed density | | `REFIT_{S3,ZMQ}_DELTA_CHANGE` | Global changed and total element counts | -| `REFIT_RECEIVER_TIMING` | Receiver batches, apply time, and verification counts | +| `REFIT_RECEIVER_TIMING` | Receiver staging span/wait, batches, apply time, and verification counts | | `REFIT_{S3,ZMQ}_DELTA_VERIFY` | Sampled transmitted-delta accuracy | | `REFIT_{S3,ZMQ}_GLOBAL_COMMIT` | Successful transfer flush | `total_s` is producer wall time. Stage fields such as `encode_s`, `s3_put_s`, and `zmq_send_s` are sums across concurrent tasks and can exceed `total_s`; do -not add them as serial phases. +not add them as serial phases. Receiver responses additionally expose node +decode/staging, worker deserialization, CPU partition, and sparse apply time. +These are also concurrent sums; compare them with receiver wall time rather +than adding them. The `partition` field is `none` for uniquely owned +policy-local shards, `names` for stable name-sharded residual exports, and +`chunks` for the full-HF compatibility path. + +Benchmark and profiler reports are dated artifacts under `profiles/`. They must +record the exact commit, image, topology, model revision, changed density, +payload settings, per-stage overlap, and sampled correctness. Do not copy an +older report's fitted values into this document as current results. The synchronizer returns metrics under `refit/delta/*`, `refit/delta_verify/*`, and `refit/transfer/*` when GRPO logs them. These are @@ -253,6 +371,18 @@ locations, and commit exact source bits only after global success. Receiver changes must preserve FIFO application, bounded memory, deferred-error propagation, flush, CUDA synchronization, and clean shutdown. +Every conversion task must remain represented by a unique local baseline unless +the complete policy-local path is disabled for FP8 parameters, quantization, or +custom FSDP. Those configurations retain the canonical full-HF baseline path. +Direct payload mappings must be rectangular affine shards of the exact HF +tensor. Tests must cover column and row offsets, gated splits, replicated +ownership, nonzero TP/ETP ranks, EP-global expert naming, DP/CP and expert-DP +ownership, transactional baseline updates, global changed-task agreement, +stable residual ownership under filtering, grouped-task expansion, and fallback +for transformed mappings. Do not infer an unknown mapping from its parameter +suffix, drop Bridge task dependencies, or modify Megatron Bridge to expose a +transport-specific hook. + Run the focused suite: ```bash @@ -285,13 +415,15 @@ the receiver before retrying. benchmark-specific estimator for the current S3 and ZeroMQ implementation. It is not a general fabric or topology model. -The sparse side embeds July 2026 end-to-end latency fits from 32 GB300 sender -GPUs in `us-east-2` to 64 H100 receiver GPUs in `us-east-1`. Measurements cover -63.2-1121.0 GB of indexed BF16 weights, S3 and ZeroMQ, raw and zstd payloads, -and 3% and 5% changed density. Any positive `--changed-pct` is accepted; values -outside 3-5% use an extrapolated power curve through the two measured-density -fits. Estimated wire bytes use the measured raw or zstd payload ratio. -The coefficients in `_SPARSE_LATENCY_FITS` implement +The zstd side embeds the July 12, 2026 end-to-end latency fits from 8 GB300 +sender GPUs in `us-east-2` to 64 H100 receiver GPUs in `us-east-1` shown above. +Measurements cover 63.2-1121.0 GB of indexed BF16 weights, S3 and ZeroMQ, and +3% and 5% changed density. The raw rows retain historical uncompressed +benchmark coefficients and are not a current production arm. Any positive +`--changed-pct` is accepted; values outside 3-5% use an extrapolated power +curve through the two measured-density fits. Estimated wire bytes use the +measured raw or zstd payload ratio. The coefficients in +`_SPARSE_LATENCY_FITS` implement `fixed_seconds + seconds_per_1000_GB * model_size_GB / 1000`; they are latency regressions, not bandwidth measurements. diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index b09795f73a2..d716bc77102 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -907,7 +907,7 @@ def _spinup_nemo_gym(base_urls, model_name): # vllm model loading prefers clean environment, initialize policy_generation before policy in colocated mode backend = generation_config["backend"] generation_config["model_name"] = policy_config["model_name"] # Needed for vLLM - refit_transport = None + remote_transport = None # Dictionary to store worker initialization timing stats for logging worker_init_timing_metrics = {} @@ -1007,6 +1007,8 @@ def init_megatron_generation(policy=None): ) return mg, time.perf_counter() - t0 + init_policy_for_generation = init_policy + def initialize_generation_with_policy( init_generation_fn, generation_name: str, @@ -1040,7 +1042,7 @@ def initialize_generation_with_policy( parallel_start_time = time.perf_counter() with ThreadPoolExecutor(max_workers=2) as executor: generation_future = executor.submit(init_generation_fn) - policy_future = executor.submit(init_policy) + policy_future = executor.submit(init_policy_for_generation) policy_generation, generation_time = generation_future.result() policy, policy_time = policy_future.result() parallel_wall_time = time.perf_counter() - parallel_start_time @@ -1062,7 +1064,7 @@ def initialize_generation_with_policy( policy_generation, generation_time = init_generation_fn() worker_init_timing_metrics[init_time_key] = generation_time - policy, policy_time = init_policy() + policy, policy_time = init_policy_for_generation() worker_init_timing_metrics["policy_init_time_s"] = policy_time worker_init_timing_metrics["parallel_init_enabled"] = 0.0 @@ -1096,6 +1098,7 @@ def initialize_generation_with_policy( elif backend == "vllm": # vLLM generation: setup config, then initialize with policy generation_config = cast(VllmConfig, generation_config) + remote_baseline_init_refs: list[Any] = [] if generation_config.get("refit_transport") is not None: # Keep optional remote transport dependencies off the default path. from nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer import ( @@ -1103,12 +1106,24 @@ def initialize_generation_with_policy( validate_vllm_remote_sparse_refit, ) - refit_transport = validate_vllm_remote_sparse_refit( - generation_config, - colocated=colocated_inference, - megatron_enabled=policy_config["megatron_cfg"]["enabled"], + remote_transport = cast( + str, + validate_vllm_remote_sparse_refit( + generation_config, + colocated=colocated_inference, + megatron_enabled=policy_config["megatron_cfg"]["enabled"], + ), ) + def init_policy_for_generation(): + policy, policy_time = init_policy() + remote_baseline_init_refs.extend( + VllmRemoteSparseWeightSynchronizer.start_baseline( + policy, remote_transport + ) + ) + return policy, policy_time + if generation_config["vllm_cfg"]["precision"] == "fp8": assert loss_config.use_importance_sampling_correction, ( "Importance sampling must be enabled for vLLM FP8 generation for good convergence!" @@ -1174,13 +1189,13 @@ def init_nemo_gym(): def init_vllm_then_policy(): pg, vllm_t = init_vllm_deferred() - p, policy_t = init_policy() + p, policy_t = init_policy_for_generation() return pg, vllm_t, p, policy_t init_tasks["vllm_policy"] = init_vllm_then_policy else: init_tasks["vllm"] = init_vllm_deferred - init_tasks["policy"] = init_policy + init_tasks["policy"] = init_policy_for_generation init_tasks["nemo_gym"] = init_nemo_gym print( @@ -1246,7 +1261,7 @@ def init_vllm_then_policy(): policy.print_node_ip_and_gpu_id() # if it is not colocated inference, initialize collective communication for update weights - if not colocated_inference and refit_transport is None: + if not colocated_inference and remote_transport is None: t0 = time.perf_counter() ip, port = train_cluster.get_master_address_and_port() print(f"Using ip: {ip}, port: {port} for collective communication", flush=True) @@ -1285,19 +1300,20 @@ def init_vllm_then_policy(): ray.get(futures_train + futures_inference) worker_init_timing_metrics["collective_init_time_s"] = time.perf_counter() - t0 - if refit_transport is not None: + if remote_transport is not None: t0 = time.perf_counter() assert isinstance(policy_generation, VllmGeneration) policy_generation.weight_synchronizer = VllmRemoteSparseWeightSynchronizer( policy, policy_generation, - transport=refit_transport.removeprefix("vllm_").removesuffix("_sparse"), + transport=remote_transport, api_key_env_var=generation_config["vllm_cfg"].get( "http_refit_api_key_env_var" ), + baseline_init_refs=remote_baseline_init_refs, ) policy_generation.weight_synchronizer.init_communicator() - worker_init_timing_metrics[f"{refit_transport}_init_time_s"] = ( + worker_init_timing_metrics[f"vllm_{remote_transport}_sparse_init_time_s"] = ( time.perf_counter() - t0 ) else: diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index c641fa74471..3e352d5fa2e 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -173,6 +173,12 @@ def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: """ 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 + def prepare_sparse_delta_refit_info( + self, state_dict_info: dict[str, tuple[tuple[int, ...], torch.dtype]] + ) -> None: + """Compile sparse placement plans before the first timed refit.""" + self._get_sparse_delta_applier().prewarm(state_dict_info) + def _maybe_process_fp8_kv_cache(self) -> None: """Process weights after loading for FP8 KV cache (static scales).""" use_fp8_kv_cache = False @@ -504,22 +510,22 @@ def update_weights_from_collective(self) -> bool: torch.cuda.empty_cache() return True - def update_weights_from_serialized_sparse_payload( + def update_weights_from_decoded_sparse_payload( self, *serialized_payloads: bytes, ) -> dict[str, Any]: - return self._get_sparse_delta_applier().update_weights_from_serialized_sparse_payload( - *serialized_payloads + return ( + self._get_sparse_delta_applier().update_weights_from_decoded_sparse_payload( + *serialized_payloads + ) ) - def update_weights_from_sparse_payload_files( + def update_weights_from_decoded_sparse_payload_files( self, *payload_paths: str, ) -> dict[str, Any]: - return ( - self._get_sparse_delta_applier().update_weights_from_sparse_payload_files( - *payload_paths - ) + return self._get_sparse_delta_applier().update_weights_from_decoded_sparse_payload_files( + *payload_paths ) def synchronize_device(self) -> None: diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py index 35432bd8d4d..f236e32d4b8 100644 --- a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py +++ b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py @@ -18,6 +18,7 @@ import re import time from collections import defaultdict +from collections.abc import Mapping from dataclasses import dataclass from math import prod from typing import Any, cast @@ -87,7 +88,7 @@ def _target_for(self, view: torch.Tensor) -> torch.Tensor | None: def __torch_dispatch__( self, func: Any, - types: Any, + _types: Any, args: tuple[Any, ...] = (), kwargs: dict[str, Any] | None = None, ) -> Any: @@ -194,35 +195,59 @@ def _compile_plans(self, metadata: list[dict[str, Any]]) -> None: copies, log_transform, identity ) + def sparse_delta_source_plans( + self, metadata: list[dict[str, Any]] + ) -> dict[str, sparse_codec.SparseSourcePlan]: + """Describe which canonical source views this worker consumes.""" + self._compile_plans(metadata) + return { + name: sparse_codec.SparseSourcePlan( + routes=tuple( + sparse_codec.SparseSourceRoute( + copy.source_offset, + copy.source_strides, + copy.shape, + copy.linear, + ) + for copy in plan.copies + ), + identity=plan.identity, + ) + for name, plan in ( + (str(item["name"]), self._plan_cache[str(item["name"])]) + for item in metadata + ) + } + + def prewarm( + self, state_dict_info: Mapping[str, tuple[tuple[int, ...], torch.dtype]] + ) -> None: + self._compile_plans( + [ + { + "name": name, + "shape": shape, + "dtype": str(dtype).removeprefix("torch."), + } + for name, (shape, dtype) in state_dict_info.items() + ] + ) + @staticmethod def _map_copy( locations: torch.Tensor, values: torch.Tensor, copy: _SparseDeltaCopyPlan, ) -> tuple[torch.Tensor, torch.Tensor]: - if copy.linear: - end = copy.source_offset + prod(copy.shape) - keep = (locations >= copy.source_offset) & (locations < end) - return ( - locations[keep] + copy.target_offset - copy.source_offset, - values[keep], - ) - mapped = torch.full_like(locations, copy.target_offset) - reconstructed = torch.full_like(locations, copy.source_offset) - relative = locations - copy.source_offset - for size, source_stride, target_stride in zip( - copy.shape, copy.source_strides, copy.target_strides, strict=True - ): - coordinate = ( - torch.div(relative, source_stride, rounding_mode="floor").remainder( - size - ) - if size > 1 - else torch.zeros_like(locations) - ) - reconstructed.add_(coordinate * source_stride) - mapped.add_(coordinate * target_stride) - keep = reconstructed == locations + mapped, keep = sparse_codec.map_sparse_locations( + locations, + copy.source_offset, + copy.source_strides, + copy.shape, + copy.linear, + copy.target_offset, + copy.target_strides, + ) return mapped[keep], values[keep] def _record_verification( @@ -307,12 +332,12 @@ def _overwrite_target_values( ) return target.data.view(-1), replacement - def _apply_item( + def _apply_decoded_item( self, item: dict[str, Any], plan: _SparseDeltaTargetPlan, - raw_locations: torch.Tensor, - raw_value_groups: tuple[torch.Tensor, ...], + locations: torch.Tensor, + values: torch.Tensor, ) -> None: operation = sparse_codec.sparse_operation(item["operation"]) source_dtype = sparse_codec.dtype_from_name(str(item["dtype"])) @@ -330,8 +355,6 @@ def _apply_item( ) if operation == "xor" and self._xor_target_mappings_overlap(plan): raise RuntimeError(f"XOR target mappings overlap for {item['name']!r}.") - value_start, value_end = int(item["value_start"]), int(item["value_end"]) - values = raw_value_groups[int(item["value_group"])][value_start:value_end] expected_dtype = sparse_codec.integer_dtype_for_element_size( source_dtype.itemsize ) @@ -344,20 +367,20 @@ def _apply_item( if plan.identity and item["index_encoding"] == "range": if operation == "xor": target = self._integer_flat(first_target) - target.narrow( - 0, int(item["range_start"]), value_end - value_start - ).bitwise_xor_(values) + target.narrow(0, int(item["range_start"]), values.numel()).bitwise_xor_( + values + ) else: target, replacement = self._overwrite_target_values( first_target, values, source_dtype, log_transform=False ) - target.narrow( - 0, int(item["range_start"]), value_end - value_start - ).copy_(replacement) + target.narrow(0, int(item["range_start"]), values.numel()).copy_( + replacement + ) return - locations = sparse_codec.sparse_locations_for_item( - item, raw_locations, device=first_target.device + locations = locations.to( + device=first_target.device, dtype=torch.int64, non_blocking=True ) if plan.identity: if operation == "xor": @@ -387,70 +410,76 @@ def _apply_item( ) target.index_copy_(0, mapped, replacement) - def _apply_sparse_weight_deltas( - self, - payload_tensors: tuple[torch.Tensor, tuple[torch.Tensor, ...]], - metadata: list[dict[str, Any]], + def _apply_decoded_sparse_weight_deltas( + self, decoded: list[sparse_codec.DecodedSparseItem] ) -> None: - self._compile_plans(metadata) - raw_locations, raw_value_groups = payload_tensors with torch.no_grad(): - for item in metadata: - self._apply_item( + for item, locations, values in decoded: + self._apply_decoded_item( item, self._plan_cache[str(item["name"])], - raw_locations, - raw_value_groups, + locations, + values, ) @wrap_with_nvtx_name( - "vllm_internal_worker_extension/update_weights_from_serialized_sparse_payload" + "vllm_internal_worker_extension/update_weights_from_decoded_sparse_payload" ) - def update_weights_from_serialized_sparse_payload( + def update_weights_from_decoded_sparse_payload( self, *serialized_payloads: bytes ) -> dict[str, Any]: - return self._load_and_apply_sparse_payloads( + return self._load_decoded_sparse_payloads( tuple(io.BytesIO(payload) for payload in serialized_payloads) ) - def _load_and_apply_sparse_payloads( + def _load_decoded_sparse_payloads( self, sources: tuple[str | io.BytesIO, ...] ) -> dict[str, Any]: started = time.perf_counter() - deserialize_s = sparse_apply_s = 0.0 - payloads = [] + deserialize_s = 0.0 + payloads: list[sparse_codec.DecodedSparsePayload] = [] for source in sources: item_started = time.perf_counter() payloads.append( cast( - sparse_codec.TensorPayload, - torch.load(source, map_location="cpu", weights_only=True), + sparse_codec.DecodedSparsePayload, + torch.load( + source, + map_location="cpu", + weights_only=True, + mmap=isinstance(source, str), + ), ) ) deserialize_s += time.perf_counter() - item_started item_started = time.perf_counter() - self._compile_plans([item for _, _, metadata in payloads for item in metadata]) + metadata = [item for _, _, items in payloads for item in items] + source_plans = self.sparse_delta_source_plans(metadata) plan_s = time.perf_counter() - item_started - for locations, values, metadata in payloads: + partition_s = sparse_apply_s = 0.0 + for payload in payloads: + item_started = time.perf_counter() + selected = sparse_codec.partition_decoded_sparse_entries( + sparse_codec.iter_decoded_sparse_payload(payload), source_plans + ) + partition_s += time.perf_counter() - item_started item_started = time.perf_counter() - self._apply_sparse_weight_deltas((locations, values), metadata) + self._apply_decoded_sparse_weight_deltas(selected) sparse_apply_s += time.perf_counter() - item_started return { "ok": True, "receiver_deserialize_s": deserialize_s, "receiver_plan_s": plan_s, + "receiver_partition_s": partition_s, "receiver_sparse_apply_s": sparse_apply_s, "receiver_total_s": time.perf_counter() - started, } - @wrap_with_nvtx_name( - "vllm_internal_worker_extension/update_weights_from_sparse_payload_files" - ) - def update_weights_from_sparse_payload_files( + def update_weights_from_decoded_sparse_payload_files( self, *payload_paths: str ) -> dict[str, Any]: - return self._load_and_apply_sparse_payloads(payload_paths) + return self._load_decoded_sparse_payloads(payload_paths) def synchronize_device(self) -> None: if torch.cuda.is_available(): @@ -461,51 +490,50 @@ def finish_sparse_delta_refit(self) -> dict[str, Any]: self.synchronize_device() verification, self._verification = self._verification, [] candidates, self._verification_candidates = self._verification_candidates, 0 - if not verification: - return { - "ok": True, - "verification_candidates": candidates, - "verification_samples": 0, - "verification_exact_mismatches": 0, - "verification_mismatches": 0, - "verification_abs_sum": 0.0, - "verification_max_abs": 0.0, - } - - with torch.no_grad(): - differences = [] - exact_mismatches = [] - mismatches = [] - samples = 0 - for target, locations, expected_bits in verification: - integer_dtype = sparse_codec.integer_dtype_for_element_size( - target.element_size() - ) - actual_bits = ( - target.data.view(integer_dtype).view(-1).index_select(0, locations) - ) - bit_mismatches = actual_bits.ne(expected_bits) - actual = target.data.view(-1).index_select(0, locations).float() - expected = expected_bits.view(target.dtype).float() - difference = torch.where( - bit_mismatches, (actual - expected).abs(), torch.zeros_like(actual) - ) - differences.append(torch.nan_to_num(difference, nan=float("inf"))) - exact_mismatches.append(bit_mismatches) - mismatches.append( - bit_mismatches - & ~torch.isclose(actual, expected, rtol=1e-6, atol=1e-8) - ) - samples += actual.numel() - difference = torch.cat(differences) - stats = torch.stack( - ( - difference.sum(), - difference.max(), - torch.cat(exact_mismatches).sum().float(), - torch.cat(mismatches).sum().float(), + samples = 0 + stats = [0.0] * 4 + if verification: + with torch.no_grad(): + differences = [] + exact_mismatches = [] + mismatches = [] + for target, locations, expected_bits in verification: + integer_dtype = sparse_codec.integer_dtype_for_element_size( + target.element_size() + ) + actual_bits = ( + target.data.view(integer_dtype) + .view(-1) + .index_select(0, locations) + ) + bit_mismatches = actual_bits.ne(expected_bits) + actual = target.data.view(-1).index_select(0, locations).float() + expected = expected_bits.view(target.dtype).float() + difference = torch.where( + bit_mismatches, + (actual - expected).abs(), + torch.zeros_like(actual), + ) + differences.append(torch.nan_to_num(difference, nan=float("inf"))) + exact_mismatches.append(bit_mismatches) + mismatches.append( + bit_mismatches + & ~torch.isclose(actual, expected, rtol=1e-6, atol=1e-8) + ) + samples += actual.numel() + difference = torch.cat(differences) + stats = ( + torch.stack( + ( + difference.sum(), + difference.max(), + torch.cat(exact_mismatches).sum().float(), + torch.cat(mismatches).sum().float(), + ) + ) + .cpu() + .tolist() ) - ).cpu() return { "ok": True, "verification_candidates": candidates, diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py index 68d35817bea..41ad65cc754 100644 --- a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py +++ b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py @@ -15,13 +15,15 @@ """Remote sparse-refit receiver lifecycle for vLLM generation workers.""" import asyncio +import io import os import tempfile import threading import time from concurrent.futures import Future, ThreadPoolExecutor -from typing import Any, Literal, cast +from typing import Any, Literal, NamedTuple, cast +import torch import uvicorn from fastapi import FastAPI, Request from fastapi.responses import JSONResponse @@ -32,9 +34,11 @@ _get_free_port_local, _get_node_ip_local, ) +from nemo_rl.utils import weight_transfer_sparse_codec as sparse_codec from nemo_rl.utils.weight_transfer_remote_sparse import ( G_VLLM_REFIT_API_KEY_HEADER, G_VLLM_REFIT_FLUSH_PATH, + G_VLLM_REFIT_PREPARE_PATH, G_VLLM_REFIT_S3_MANIFEST_PATH, decode_sparse_payload, download_s3_refit_payload, @@ -52,6 +56,56 @@ ) +def _decode_staged_payload( + serialized: bytes, +) -> tuple[sparse_codec.DecodedSparsePayload, int]: + payload = cast( + sparse_codec.TensorPayload, + torch.load(io.BytesIO(serialized), map_location="cpu", weights_only=True), + ) + return ( + sparse_codec.decode_sparse_tensor_payload_for_staging(payload), + sum(len(item.get("verification_locations", ())) for item in payload[2]), + ) + + +class _StagedSparsePayload(NamedTuple): + path: str + started_at: float + finished_at: float + deserialize_s: float + save_s: float + candidates: int + + +def _stage_sparse_payload( + serialized: bytes, + staging_dir: str, +) -> _StagedSparsePayload: + started_at = time.perf_counter() + decoded, candidates = _decode_staged_payload(serialized) + deserialize_s = time.perf_counter() - started_at + descriptor, path = tempfile.mkstemp( + prefix="nemo_rl_refit_", suffix=".pt", dir=staging_dir + ) + os.close(descriptor) + started = time.perf_counter() + try: + torch.save(decoded, path) + except Exception: + os.unlink(path) + raise + finished_at = time.perf_counter() + return _StagedSparsePayload( + path, + started_at, + finished_at, + deserialize_s, + finished_at - started, + candidates, + ) + + class VllmSparseRefitReceiver: """Own the optional transport server, apply queue, and relay resources.""" @@ -63,15 +117,25 @@ def __init__(self, worker: Any) -> None: thread_name_prefix="nrl-vllm-sparse-refit", ) self._refit_apply_futures: list[Future[dict[str, Any]]] = [] - self._refit_apply_pending_payloads: list[bytes] = [] + self._refit_apply_pending_payloads: list[ + bytes | Future[_StagedSparsePayload] + ] = [] self._refit_seen_payloads: dict[tuple[str, int, int], str] = {} self._refit_workers_share_node = False self._refit_apply_queue_depth = refit_env_int( - "NRL_REFIT_APPLY_QUEUE_DEPTH", default=2 + "NRL_REFIT_APPLY_QUEUE_DEPTH", default=32 ) self._refit_apply_batch_size = refit_env_int( "NRL_REFIT_APPLY_BATCH_SIZE", default=8 ) + self._refit_partition_executor = ThreadPoolExecutor( + max_workers=refit_env_int( + "NRL_REFIT_PARTITION_WORKERS", + default=max(2, min(8, os.cpu_count() or 8)), + ), + thread_name_prefix="nrl-vllm-sparse-partition", + ) + self._refit_verification_candidates = 0 self._refit_batch_staging_dir = ( os.getenv("NRL_REFIT_BATCH_STAGING_DIR") or "/dev/shm" ) @@ -96,6 +160,7 @@ def shutdown(self) -> None: self._flush_queued_sparse_payloads() self._refit_apply_executor.shutdown(wait=True) + self._refit_partition_executor.shutdown(wait=True) if self._refit_http_server is not None: self._refit_http_server[1].join(timeout=5.0) @@ -108,7 +173,6 @@ def _enqueue_sparse_payload_apply( checksum: str, ) -> dict[str, Any]: completed: list[Future[dict[str, Any]]] = [] - submitted = None with self._refit_apply_queue_condition: seen_checksum = self._refit_seen_payloads.get(payload_key) if seen_checksum is not None: @@ -126,21 +190,33 @@ def _enqueue_sparse_payload_apply( completed.append(self._refit_apply_futures.pop(0)) response = self._collect_refit_apply_results(completed) self._refit_seen_payloads[payload_key] = checksum - self._refit_apply_pending_payloads.append(payload) + pending: bytes | Future[_StagedSparsePayload] = payload + if self._refit_workers_share_node: + pending = self._refit_partition_executor.submit( + _stage_sparse_payload, + payload, + self._refit_batch_staging_dir, + ) + self._refit_apply_pending_payloads.append(pending) if len(self._refit_apply_pending_payloads) == self._refit_apply_batch_size: - submitted = self._submit_pending_sparse_payloads() - if submitted is not None: - submitted.add_done_callback(self._notify_refit_apply_waiters) + self._submit_pending_sparse_payloads() return response def _submit_pending_sparse_payloads(self) -> Future[dict[str, Any]]: payloads = tuple(self._refit_apply_pending_payloads) self._refit_apply_pending_payloads.clear() - future = self._refit_apply_executor.submit( - self.update_weights_from_serialized_sparse_payloads, - payloads, - ) + if self._refit_workers_share_node: + future = self._refit_apply_executor.submit( + self.update_weights_from_staged_sparse_payloads, + cast(tuple[Future[_StagedSparsePayload], ...], payloads), + ) + else: + future = self._refit_apply_executor.submit( + self.update_weights_from_serialized_sparse_payloads, + cast(tuple[bytes, ...], payloads), + ) self._refit_apply_futures.append(future) + future.add_done_callback(self._notify_refit_apply_waiters) return future def _notify_refit_apply_waiters(self, _future: Future[dict[str, Any]]) -> None: @@ -192,45 +268,85 @@ def update_weights_from_serialized_sparse_payloads( serialized_payloads: tuple[bytes, ...], ) -> dict[str, Any]: """Apply a FIFO batch of sparse deltas through one collective RPC.""" - if not self._refit_workers_share_node: - response = self._refit_collective_response( - self._refit_collective_rpc( - "update_weights_from_serialized_sparse_payload", - serialized_payloads, - ) + + def decode_for_rpc(serialized: bytes) -> tuple[bytes, int]: + decoded, candidates = _decode_staged_payload(serialized) + buffer = io.BytesIO() + torch.save(decoded, buffer) + return buffer.getvalue(), candidates + + decoded_payloads = list( + self._refit_partition_executor.map(decode_for_rpc, serialized_payloads) + ) + self._refit_verification_candidates += sum( + candidates for _, candidates in decoded_payloads + ) + response = self._refit_collective_response( + self._refit_collective_rpc( + "update_weights_from_decoded_sparse_payload", + tuple(payload for payload, _ in decoded_payloads), + ) + ) + response["payloads"] = len(serialized_payloads) + return response + + def update_weights_from_staged_sparse_payloads( + self, + staged_payloads: tuple[Future[_StagedSparsePayload], ...], + ) -> dict[str, Any]: + started = time.perf_counter() + staged: list[_StagedSparsePayload] = [] + try: + stage_error = None + for future in staged_payloads: + try: + staged.append(future.result()) + except Exception as exc: + stage_error = stage_error or exc + if stage_error is not None: + raise stage_error + stage_wait_s = time.perf_counter() - started + self._refit_verification_candidates += sum( + payload.candidates for payload in staged ) - response["payloads"] = len(serialized_payloads) - return response - - with tempfile.TemporaryDirectory( - prefix="nemo_rl_refit_", dir=self._refit_batch_staging_dir - ) as staging_dir: - paths = [] - for index, payload in enumerate(serialized_payloads): - path = os.path.join(staging_dir, str(index)) - with open(path, "wb") as staged: - staged.write(payload) - paths.append(path) try: response = self._refit_collective_response( self._refit_collective_rpc( - "update_weights_from_sparse_payload_files", - tuple(paths), + "update_weights_from_decoded_sparse_payload_files", + tuple(payload.path for payload in staged), ) ) except Exception: - # Drain peers before TemporaryDirectory removes shared batch files. + # Drain peers before removing shared batch files. self._refit_collective_rpc("synchronize_device", ()) raise - response["payloads"] = len(serialized_payloads) + finally: + for payload in staged: + os.unlink(payload.path) + worker_total_s = float(response.get("receiver_total_s", 0.0)) + response.update( + receiver_node_deserialize_s=max( + (payload.deserialize_s for payload in staged), default=0.0 + ), + receiver_stage_s=( + max(payload.finished_at for payload in staged) + - min(payload.started_at for payload in staged) + ), + receiver_stage_save_s=max( + (payload.save_s for payload in staged), default=0.0 + ), + receiver_stage_wait_s=stage_wait_s, + ) + response["receiver_worker_total_s"] = worker_total_s + response["receiver_total_s"] = time.perf_counter() - started + response["payloads"] = len(staged_payloads) return response def _flush_queued_sparse_payloads(self) -> dict[str, Any]: started = time.perf_counter() - submitted = None with self._refit_apply_queue_condition: if self._refit_apply_pending_payloads: - submitted = self._submit_pending_sparse_payloads() + self._submit_pending_sparse_payloads() futures = list(self._refit_apply_futures) self._refit_apply_futures.clear() self._refit_apply_queue_condition.notify_all() @@ -238,17 +354,19 @@ def _flush_queued_sparse_payloads(self) -> dict[str, Any]: batch_count = ( payload_count + self._refit_apply_batch_size - 1 ) // self._refit_apply_batch_size - if submitted is not None: - submitted.add_done_callback(self._notify_refit_apply_waiters) response = self._collect_refit_apply_results(futures) if futures: - response.update( - self._refit_collective_response( - self._refit_collective_rpc("finish_sparse_delta_refit", ()) - ) + verification = self._refit_collective_response( + self._refit_collective_rpc("finish_sparse_delta_refit", ()) ) + if self._refit_verification_candidates: + verification["verification_candidates"] = ( + self._refit_verification_candidates + ) + response.update(verification) with self._refit_apply_queue_condition: self._refit_seen_payloads.clear() + self._refit_verification_candidates = 0 response.update( payloads=payload_count, batches=batch_count, @@ -273,6 +391,23 @@ def _flush_queued_sparse_payloads(self) -> dict[str, Any]: ) return response + def _prepare_sparse_refit_info(self, request: dict[str, Any]) -> dict[str, Any]: + started = time.perf_counter() + state_dict_info = { + name: (tuple(shape), sparse_codec.dtype_from_name(dtype)) + for name, (shape, dtype) in request["tensors"].items() + } + self._refit_collective_rpc( + "prepare_sparse_delta_refit_info", (state_dict_info,) + ) + seconds = time.perf_counter() - started + print( + f"REFIT_RECEIVER_PREWARM tensors={len(state_dict_info)} " + f"seconds={seconds:.3f}", + flush=True, + ) + return {"ok": True, "tensors": len(state_dict_info), "seconds": seconds} + async def _apply_s3_manifest_payload( self, manifest: dict[str, Any], @@ -322,7 +457,7 @@ def setup_api_server(self, app: Any) -> None: async def respond( raw_request: Request, - action: Literal["s3", "flush", "zmq"], + action: Literal["prepare", "s3", "flush", "zmq"], ) -> JSONResponse: if cfg["vllm_cfg"]["async_engine"]: self._refit_async_loop = asyncio.get_running_loop() @@ -334,7 +469,12 @@ async def respond( content={"ok": False, "error": "unauthorized"}, status_code=403 ) try: - if action == "s3": + if action == "prepare": + result = await asyncio.to_thread( + self._prepare_sparse_refit_info, + await raw_request.json(), + ) + elif action == "s3": result = await self._apply_s3_manifest_payload( await raw_request.json() ) @@ -349,20 +489,25 @@ async def respond( status_code=200 if result.get("ok") is True else 500, ) - @app.post(G_VLLM_REFIT_S3_MANIFEST_PATH) - async def apply_s3_manifest_refit(raw_request: Request) -> JSONResponse: - return await respond(raw_request, "s3") + def endpoint(action: Literal["prepare", "s3", "flush", "zmq"]): + async def handle(raw_request: Request) -> JSONResponse: + return await respond(raw_request, action) - @app.post(G_VLLM_REFIT_FLUSH_PATH) - async def flush_sparse_delta_refit(raw_request: Request) -> JSONResponse: - return await respond(raw_request, "flush") + return handle - @app.post(G_VLLM_REFIT_ZMQ_PAYLOAD_PATH) - async def apply_zmq_sparse_refit(raw_request: Request) -> JSONResponse: - return await respond(raw_request, "zmq") + for path, action in ( + (G_VLLM_REFIT_S3_MANIFEST_PATH, "s3"), + (G_VLLM_REFIT_PREPARE_PATH, "prepare"), + (G_VLLM_REFIT_FLUSH_PATH, "flush"), + (G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, "zmq"), + ): + app.add_api_route(path, endpoint(action), methods=["POST"]) def report_refit_server_base_url(self) -> str | None: - return self._refit_http_server[2] if self._refit_http_server else None + if self._refit_http_server is not None: + return self._refit_http_server[2] + base_url = getattr(self._worker, "base_url", None) + return base_url.removesuffix("/v1") if base_url else None def start_zmq_sparse_refit_relay(self, refit_urls: list[str]) -> str: if self._zmq_refit_server is not None: diff --git a/nemo_rl/models/generation/vllm/vllm_worker.py b/nemo_rl/models/generation/vllm/vllm_worker.py index 316ec5c4c64..e3e1791272b 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker.py +++ b/nemo_rl/models/generation/vllm/vllm_worker.py @@ -658,19 +658,15 @@ def _get_raw_spec_counters(self) -> dict[str, float | list[float]]: metrics[metric.name] = metric.value return metrics - def _require_sparse_refit_receiver(self) -> Any: - if self._sparse_refit_receiver is None: - raise RuntimeError("Remote sparse refit is not enabled for this worker.") - return self._sparse_refit_receiver - def report_refit_server_base_url(self) -> str | None: receiver = self._sparse_refit_receiver return receiver.report_refit_server_base_url() if receiver is not None else None def start_zmq_sparse_refit_relay(self, refit_urls: list[str]) -> str: - return self._require_sparse_refit_receiver().start_zmq_sparse_refit_relay( - refit_urls - ) + receiver = self._sparse_refit_receiver + if receiver is None: + raise RuntimeError("Remote sparse refit is not enabled for this worker.") + return receiver.start_zmq_sparse_refit_relay(refit_urls) def stop_zmq_sparse_refit_relay(self) -> None: receiver = self._sparse_refit_receiver diff --git a/nemo_rl/models/generation/vllm/vllm_worker_async.py b/nemo_rl/models/generation/vllm/vllm_worker_async.py index 17fcd30e489..2f50e4a150e 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker_async.py +++ b/nemo_rl/models/generation/vllm/vllm_worker_async.py @@ -442,11 +442,6 @@ async def get_reserved_url(self) -> Optional[str]: async def report_dp_openai_server_base_url(self) -> Optional[str]: return self.base_url - def report_refit_server_base_url(self) -> str | None: - if self.cfg.get("refit_transport") is None or self.base_url is None: - return None - return self.base_url.removesuffix("/v1") - # ruff: noqa def _setup_vllm_openai_api_server(self, app: FastAPI) -> FastAPI: from copy import deepcopy diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 28b49cb0eac..209e06371f7 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -153,6 +153,7 @@ class MegatronPolicyWorkerImpl( # begin/abort; None when no step is open. Declared at class level so # ``self._train_step_state = None`` after finish/abort type-checks. _train_step_state: Optional[dict[str, Any]] = None + _remote_sparse_refit: Any = None def __repr__(self): """Customizes the actor's prefix in the Ray logs. @@ -364,18 +365,6 @@ def __init__( self.is_generation_colocated = runtime_config.is_generation_colocated self.final_padded_vocab_size = runtime_config.final_padded_vocab_size self.sampling_params = runtime_config.sampling_params - generation_config = self.cfg.get("generation") - delta_config = None - if generation_config and generation_config.get("refit_transport") is not None: - delta_config = cast(VllmConfig, generation_config).get("delta_compression") - self._remote_sparse_refit = None - if delta_config: - # Keep codec and remote transport state out of standard policy workers. - from nemo_rl.models.policy.workers.megatron_remote_sparse_refit import ( - MegatronRemoteSparseRefit, - ) - - self._remote_sparse_refit = MegatronRemoteSparseRefit(self, delta_config) self.defer_fp32_logits = self.cfg["megatron_cfg"].get( "defer_fp32_logits", None @@ -1815,8 +1804,8 @@ def init_remote_sparse_delta_baseline( shard_rank: int, shard_count: int, transport: str, - ) -> None: - self._require_remote_sparse_refit().initialize_baseline( + ) -> dict[str, tuple[tuple[int, ...], torch.dtype]]: + return self._require_remote_sparse_refit().initialize_baseline( shard_rank=shard_rank, shard_count=shard_count, transport=transport, @@ -1847,14 +1836,12 @@ def stream_remote_sparse_weights( def _require_remote_sparse_refit(self) -> Any: if self._remote_sparse_refit is None: - raise RuntimeError("Remote sparse refit is not enabled for this worker.") - return self._remote_sparse_refit + from nemo_rl.models.policy.workers.megatron_remote_sparse_refit import ( + MegatronRemoteSparseRefit, + ) - def _get_refit_conversion_tasks(self) -> list[Any]: - if self.refit_conversion_tasks is None: - tasks = self.megatron_bridge.get_conversion_tasks([self.model]) - self.refit_conversion_tasks = [task for task in tasks if task is not None] - return self.refit_conversion_tasks + self._remote_sparse_refit = MegatronRemoteSparseRefit.from_worker(self) + return self._remote_sparse_refit def finish_remote_sparse_delta_sync(self, *, succeeded: bool) -> None: self._require_remote_sparse_refit().finish(succeeded) @@ -1873,7 +1860,11 @@ def _calculate_refit_param_info(self) -> list[tuple[str, int]]: Returns: List of (parameter_name, size_in_bytes) tuples. """ - conversion_tasks = self._get_refit_conversion_tasks() + self.refit_conversion_tasks = [ + task + for task in self.megatron_bridge.get_conversion_tasks([self.model]) + if task is not None + ] param_info = [] def calculate_size_in_bytes(param, tp_size, ep_size): @@ -1897,7 +1888,7 @@ def calculate_size_in_bytes(param, tp_size, ep_size): # Broadcast size_in_bytes across pipeline parallel ranks return broadcast_obj_from_pp_rank(size_in_bytes) - for task in conversion_tasks: + for task in self.refit_conversion_tasks: param_info.append( ( task.param_name, @@ -1913,6 +1904,7 @@ def calculate_size_in_bytes(param, tp_size, ep_size): def _iter_params_with_optional_kv_scales( self, kv_scales: Optional[dict[str, float]] = None, + conversion_tasks: Optional[list[Any]] = None, ) -> Iterator[tuple[str, torch.Tensor]]: """Yield exported HF parameters and optionally append FP8 KV/Q scale tensors. @@ -1926,7 +1918,11 @@ def _iter_params_with_optional_kv_scales( base_iter = self.megatron_bridge.export_hf_weights( [self.model], show_progress=False, - conversion_tasks=self._get_refit_conversion_tasks(), + conversion_tasks=( + self.refit_conversion_tasks + if conversion_tasks is None + else conversion_tasks + ), ) # Yield the original parameters first. diff --git a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py index b3e36deeeac..647b1100161 100644 --- a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py +++ b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py @@ -14,6 +14,10 @@ """Optional remote sparse-refit state owned by a Megatron policy worker.""" +import re +from collections.abc import Iterable, Mapping +from concurrent.futures import ThreadPoolExecutor +from functools import cache, partial from typing import Any import torch @@ -21,16 +25,377 @@ from nemo_rl.utils.weight_transfer_remote_sparse import ( SparseDeltaStreamResult, init_sparse_delta_baseline_from_iterator, + sparse_name_shard, stream_sparse_delta_payloads_via_s3_manifest, ) -from nemo_rl.utils.weight_transfer_sparse_codec import DeltaCompressionTracker +from nemo_rl.utils.weight_transfer_sparse_codec import ( + DeltaCompressionTracker, + SparseShardProjection, +) from nemo_rl.utils.weight_transfer_zmq import stream_sparse_delta_payloads_via_zmq +_UNSUPPORTED = 0 +_COLUMN = 1 +_ROW = 2 +_REPLICATED = 3 +_DIRECT = 4 +_GATED = 5 + class MegatronRemoteSparseRefit: - def __init__(self, worker: Any, delta_config: dict[str, Any]) -> None: + @classmethod + def from_worker(cls, worker: Any) -> "MegatronRemoteSparseRefit": + generation_config = worker.cfg.get("generation") or {} + delta_config = generation_config.get("delta_compression") + if generation_config.get("refit_transport") is None or not delta_config: + raise RuntimeError("Remote sparse refit is not enabled for this worker.") + return cls(worker, delta_config) + + def __init__(self, worker: Any, delta_config: Mapping[str, Any]) -> None: self._worker = worker - self._tracker = DeltaCompressionTracker(delta_config) + self._delta_config = delta_config + residual_config = dict(delta_config) + if residual_config["encoding"] == "xor": + residual_config["encoding"] = "overwrite" + self._tracker = DeltaCompressionTracker(residual_config) + self._local_tracker: DeltaCompressionTracker | None = None + self._change_tracker: DeltaCompressionTracker | None = None + self._local_tensors: list[tuple[str, torch.Tensor]] = [] + self._misc_local_tensors: list[tuple[str, torch.Tensor]] = [] + self._misc_conversion_tasks: list[Any] | None = None + self._filter_misc_tasks = False + self._uses_policy_local_path = False + + @staticmethod + @cache + def _bridge_mapping_types() -> tuple[Any, Any, Any, Any, Any, Any]: + # Bridge is optional outside Megatron workers, so keep these imports local. + from megatron.bridge.models.conversion.param_mapping import ( + AutoMapping, + ColumnParallelMapping, + DirectMapping, + GatedMLPMapping, + ReplicatedMapping, + RowParallelMapping, + ) + + return ( + AutoMapping, + ColumnParallelMapping, + DirectMapping, + GatedMLPMapping, + ReplicatedMapping, + RowParallelMapping, + ) + + @staticmethod + def _all_reduce_max(values: list[int]) -> list[int]: + if not values or not torch.distributed.is_initialized(): + return values + backend = str(torch.distributed.get_backend()).lower() + device = ( + torch.device("cuda", torch.cuda.current_device()) + if backend.endswith("nccl") + else torch.device("cpu") + ) + reduced = torch.tensor(values, dtype=torch.int32, device=device) + torch.distributed.all_reduce(reduced, op=torch.distributed.ReduceOp.MAX) + return reduced.cpu().tolist() + + @classmethod + def _local_mapping_kind(cls, task: Any) -> int: + ( + AutoMapping, + *mapping_types, + ) = cls._bridge_mapping_types() + mapping = task.mapping + for mapping_type, kind in zip( + mapping_types, + (_COLUMN, _DIRECT, _GATED, _REPLICATED, _ROW), + strict=True, + ): + if type(mapping) is mapping_type: + return kind + if ( + type(mapping) is AutoMapping + and mapping.permute_dims is None + and task.megatron_module is not None + ): + return { + "column": _COLUMN, + "row": _ROW, + "replicated": _REPLICATED, + }.get(mapping._detect_parallelism_type(task.megatron_module), _UNSUPPORTED) + return _UNSUPPORTED + + def _bridge_exports_are_identity(self) -> bool: + bridge = getattr(self._worker, "megatron_bridge", None) + if bridge is None: + return True + from megatron.bridge.models.conversion.model_bridge import MegatronModelBridge + + model_bridge = bridge._model_bridge + return ( + type(model_bridge).maybe_modify_converted_hf_weight + is MegatronModelBridge.maybe_modify_converted_hf_weight + ) + + @staticmethod + def _is_padded_or_tied_weight(task: Any) -> bool: + hf_names = ( + task.mapping.hf_param.values() + if isinstance(task.mapping.hf_param, dict) + else (task.mapping.hf_param,) + ) + return task.global_param_name.endswith( + ("embedding.word_embeddings.weight", "output_layer.weight") + ) or any( + str(name).endswith( + ("embed_tokens.weight", "embeddings.weight", "lm_head.weight") + ) + for name in hf_names + ) + + @staticmethod + def _can_project_task(task: Any, kind: int, *, identity_export: bool) -> bool: + if kind == _UNSUPPORTED or not identity_export: + return False + mapping = task.mapping + if getattr(mapping, "is_grouped_export", False) or getattr( + mapping, "is_adapter", False + ): + return False + if MegatronRemoteSparseRefit._is_padded_or_tied_weight(task): + return False + if kind == _GATED: + return isinstance(mapping.hf_param, dict) and set(mapping.hf_param) == { + "gate", + "up", + } + return isinstance(mapping.hf_param, str) + + @staticmethod + def _owns_policy_local_task(task: Any, *, replicated: bool = False) -> bool: + if not torch.distributed.is_initialized(): + return True + from megatron.core import parallel_state + + if task.mapping.is_expert: + # Expert-DP includes DP replicas and TP replicas when ETP < TP. + replica_rank = parallel_state.get_expert_data_parallel_rank() + replica_count = parallel_state.get_expert_data_parallel_world_size() + else: + replica_rank = parallel_state.get_data_parallel_rank( + with_context_parallel=True + ) + replica_count = parallel_state.get_data_parallel_world_size( + with_context_parallel=True + ) + if replicated: + replica_rank = ( + replica_rank * task.mapping.tp_size + task.mapping.tp_rank + ) + replica_count *= task.mapping.tp_size + return replica_rank == sparse_name_shard(task.global_param_name, replica_count) + + def _policy_local_path_is_safe(self) -> bool: + config = getattr(self._worker, "cfg", {}) + if config.get("quant_cfg") is not None: + return False + ddp_config = config.get("megatron_cfg", {}).get( + "distributed_data_parallel_config", {} + ) + if ddp_config.get("use_custom_fsdp", False): + return False + fp8_cfg = getattr(self._worker, "fp8_cfg", None) + return not fp8_cfg or not fp8_cfg.get("fp8_param", False) + + @staticmethod + def _canonical_hf_name(task: Any, name: str) -> str: + mapping = task.mapping + if not mapping.is_expert or mapping.ep_size == 1: + return name + + match = re.search(r"(\.experts\.)(\d+)(\.)", name) + config = getattr(task.megatron_module, "config", None) + num_experts = getattr(config, "num_moe_experts", None) + if match is None or not isinstance(num_experts, int): + raise ValueError(f"Cannot project expert parameter {name!r}.") + if num_experts % mapping.ep_size: + raise ValueError( + f"Expert count {num_experts} is not divisible by EP size " + f"{mapping.ep_size}." + ) + experts_per_rank = num_experts // mapping.ep_size + expert = int(match.group(2)) % experts_per_rank + expert += experts_per_rank * mapping.ep_rank + return f"{name[: match.start(2)]}{expert}{name[match.end(2) :]}" + + @staticmethod + def _projection( + name: str, + tensor: torch.Tensor, + *, + shard_dim: int, + shard_rank: int, + shard_count: int, + ) -> SparseShardProjection: + if tensor.ndim <= shard_dim: + raise ValueError(f"Cannot shard {name!r} on dimension {shard_dim}.") + global_shape = list(tensor.shape) + offsets = [0] * tensor.ndim + global_shape[shard_dim] *= shard_count + offsets[shard_dim] = tensor.shape[shard_dim] * shard_rank + return SparseShardProjection(name, tuple(global_shape), tuple(offsets)) + + @classmethod + def _task_local_tensors( + cls, task: Any, kind: int + ) -> list[tuple[str, torch.Tensor, SparseShardProjection]]: + if task.param_weight is None: + return [] + + mapping = task.mapping + tensor = task.param_weight + replicated = kind in (_DIRECT, _REPLICATED) or ( + kind == _ROW and tensor.ndim == 1 + ) + if not cls._owns_policy_local_task(task, replicated=replicated): + return [] + if kind == _GATED: + gate, up = torch.chunk(tensor, 2, dim=0) + return [ + ( + f"{task.global_param_name}:{role}", + value, + cls._projection( + cls._canonical_hf_name(task, str(mapping.hf_param[role])), + value, + shard_dim=0, + shard_rank=mapping.tp_rank, + shard_count=mapping.tp_size, + ), + ) + for role, value in (("gate", gate), ("up", up)) + ] + + name = cls._canonical_hf_name(task, str(mapping.hf_param)) + projection = ( + SparseShardProjection(name, tuple(tensor.shape), (0,) * tensor.ndim) + if replicated + else cls._projection( + name, + tensor, + shard_dim=0 if kind == _COLUMN else 1, + shard_rank=mapping.tp_rank, + shard_count=mapping.tp_size, + ) + ) + return [(task.global_param_name, tensor, projection)] + + def _prepare_paths(self) -> None: + if self._misc_conversion_tasks is not None: + return + tasks = [ + task + for task in self._worker.megatron_bridge.get_conversion_tasks( + [self._worker.model] + ) + if task is not None + ] + misc_tasks: list[Any] = [] + self._misc_conversion_tasks = misc_tasks + if not tasks or not self._policy_local_path_is_safe(): + misc_tasks.extend(tasks) + return + self._uses_policy_local_path = True + kinds = self._all_reduce_max([self._local_mapping_kind(task) for task in tasks]) + identity_export = self._bridge_exports_are_identity() + self._filter_misc_tasks = identity_export + projections = {} + for task, kind in zip(tasks, kinds, strict=True): + if self._can_project_task(task, kind, identity_export=identity_export): + for key, tensor, projection in self._task_local_tensors(task, kind): + if key in projections: + raise ValueError( + f"Duplicate policy-local sparse shard {key!r}." + ) + projections[key] = projection + self._local_tensors.append((key, tensor)) + continue + + task_index = len(misc_tasks) + misc_tasks.append(task) + if task.param_weight is None: + continue + replicated = kind in (_DIRECT, _REPLICATED) or ( + kind == _ROW and task.param_weight.ndim == 1 + ) + if not self._owns_policy_local_task(task, replicated=replicated): + continue + key = f"{task_index}:{task.global_param_name}" + self._misc_local_tensors.append((key, task.param_weight)) + + if projections: + self._local_tracker = DeltaCompressionTracker( + self._delta_config, projections=projections + ) + if self._misc_local_tensors: + self._change_tracker = DeltaCompressionTracker(self._delta_config) + + def _iter_misc_params( + self, conversion_tasks: list[Any] | None = None + ) -> Iterable[tuple[str, torch.Tensor]]: + return self._worker._iter_params_with_optional_kv_scales( + conversion_tasks=( + self._misc_conversion_tasks + if conversion_tasks is None + else conversion_tasks + ) + ) + + def _changed_misc_tasks(self) -> tuple[list[Any], int, int]: + assert self._misc_conversion_tasks is not None + changed_keys: set[str] = set() + changed = total = 0 + if self._change_tracker is not None: + changed_keys, changed, total = self._change_tracker.prepare_change_summary( + self._misc_local_tensors + ) + flags = [0] * len(self._misc_conversion_tasks) + for key in changed_keys: + flags[int(key.partition(":")[0])] = 1 + flags = self._all_reduce_max(flags) + if any(flags) and not self._filter_misc_tasks: + flags = [1] * len(flags) + grouped_keys = { + task.mapping.group_key + for task, task_changed in zip( + self._misc_conversion_tasks, flags, strict=True + ) + if task_changed and getattr(task.mapping, "is_grouped_export", False) + } + flags = [ + task_changed + or ( + getattr(task.mapping, "is_grouped_export", False) + and task.mapping.group_key in grouped_keys + ) + for task, task_changed in zip( + self._misc_conversion_tasks, flags, strict=True + ) + ] + return ( + [ + task + for task, task_changed in zip( + self._misc_conversion_tasks, flags, strict=True + ) + if task_changed + ], + changed, + total, + ) def initialize_baseline( self, @@ -38,14 +403,48 @@ def initialize_baseline( shard_rank: int, shard_count: int, transport: str, - ) -> None: - init_sparse_delta_baseline_from_iterator( - self._worker._iter_params_with_optional_kv_scales(), - delta_tracker=self._tracker, + ) -> dict[str, tuple[tuple[int, ...], torch.dtype]]: + self._prepare_paths() + snapshot = partial( + init_sparse_delta_baseline_from_iterator, shard_rank=shard_rank, shard_count=shard_count, transport=transport, ) + if not self._uses_policy_local_path: + snapshot(self._iter_misc_params(), delta_tracker=self._tracker) + return self.refit_info() + + def snapshot_local_baselines() -> None: + for tracker, tensors in ( + (self._local_tracker, self._local_tensors), + (self._change_tracker, self._misc_local_tensors), + ): + if tracker is not None: + snapshot(tensors, delta_tracker=tracker, partition="none") + + with ThreadPoolExecutor( + max_workers=1, thread_name_prefix="nrl-refit-policy-local" + ) as executor: + local_future = executor.submit(snapshot_local_baselines) + snapshot( + self._iter_misc_params(), + delta_tracker=self._tracker, + partition="names", + ) + local_future.result() + return self.refit_info() + + def refit_info(self) -> dict[str, tuple[tuple[int, ...], torch.dtype]]: + info = { + name: (tuple(tensor.shape), tensor.dtype) + for name, tensor in self._tracker.baseline.items() + } + if self._local_tracker is not None: + for name, tensor in self._local_tensors: + projection = self._local_tracker.projections[name] + info[projection.name] = (projection.global_shape, tensor.dtype) + return info def stream( self, @@ -58,26 +457,64 @@ def stream( shard_rank: int, shard_count: int, ) -> SparseDeltaStreamResult: + self._prepare_paths() streamer = { "s3": stream_sparse_delta_payloads_via_s3_manifest, "zmq": stream_sparse_delta_payloads_via_zmq, }[transport] - result = streamer( - self._worker._iter_params_with_optional_kv_scales(), - delta_tracker=self._tracker, + send = partial( + streamer, refit_targets=targets, - transfer_id=transfer_id, api_key_env_var=api_key_env_var, timeout_s=timeout_s, shard_rank=shard_rank, shard_count=shard_count, ) + if not self._uses_policy_local_path: + result = send( + self._iter_misc_params(), + delta_tracker=self._tracker, + transfer_id=transfer_id, + ) + else: + with ThreadPoolExecutor( + max_workers=1, thread_name_prefix="nrl-refit-policy-local" + ) as executor: + local_future = None + if self._local_tracker is not None: + local_future = executor.submit( + send, + self._local_tensors, + delta_tracker=self._local_tracker, + transfer_id=f"{transfer_id}-local", + partition="none", + ) + changed_tasks, misc_changed, misc_total = self._changed_misc_tasks() + misc_result = send( + self._iter_misc_params(changed_tasks), + delta_tracker=self._tracker, + transfer_id=f"{transfer_id}-misc", + partition="names", + ) + local_result = ( + local_future.result() + if local_future is not None + else {"payloads": 0, "changed_elements": 0, "total_elements": 0} + ) + result = SparseDeltaStreamResult( + payloads=int(local_result["payloads"]) + int(misc_result["payloads"]), + changed_elements=int(local_result["changed_elements"]) + misc_changed, + total_elements=int(local_result["total_elements"]) + misc_total, + ) if torch.cuda.is_available(): torch.cuda.synchronize() return result def finish(self, succeeded: bool) -> None: - if succeeded: - self._tracker.on_sync_succeeded() - else: - self._tracker.on_sync_failed() + for tracker in (self._tracker, self._local_tracker, self._change_tracker): + if tracker is None: + continue + if succeeded: + tracker.on_sync_succeeded() + else: + tracker.on_sync_failed() diff --git a/nemo_rl/utils/weight_transfer_remote_sparse.py b/nemo_rl/utils/weight_transfer_remote_sparse.py index b4f2c28faa9..8c084a84300 100644 --- a/nemo_rl/utils/weight_transfer_remote_sparse.py +++ b/nemo_rl/utils/weight_transfer_remote_sparse.py @@ -19,11 +19,12 @@ import os import threading import time -from collections.abc import Iterable, Iterator, Mapping, Sequence +from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait from contextlib import suppress +from dataclasses import dataclass from functools import cache -from typing import Any, Callable, TypedDict +from typing import Any, Literal, TypedDict from urllib.parse import quote import requests @@ -36,9 +37,12 @@ DeltaCompressionTracker, NamedTensor, TensorBatch, + TensorPayload, + merge_sparse_payloads, ) G_VLLM_REFIT_S3_MANIFEST_PATH = "/nemo-rl/refit/s3-manifest" +G_VLLM_REFIT_PREPARE_PATH = "/nemo-rl/refit/prepare" G_VLLM_REFIT_FLUSH_PATH = "/nemo-rl/refit/flush" G_VLLM_REFIT_API_KEY_HEADER = "x-nemo-rl-refit-key" _CONTROL_SESSION_LOCAL = threading.local() @@ -52,6 +56,17 @@ class SparseDeltaStreamResult(TypedDict): total_elements: int +SparsePartitionMode = Literal["chunks", "names", "none"] + + +@dataclass +class _SparsePayloadBucket: + payloads: list[TensorPayload] + dense_bytes: int = 0 + encode_s: float = 0.0 + next_index: int = 0 + + @cache def _s3_client(region: str) -> Any: # Keep the AWS runtime unloaded for ZeroMQ-only jobs. @@ -93,7 +108,7 @@ def put(self, key: str, body: bytes) -> None: request=self._request("PUT", key, body), ).finished_future.result() - def get(self, key: str) -> bytes: + def get(self, key: str) -> bytearray: from awscrt.http import HttpHeaders body = bytearray() @@ -120,7 +135,7 @@ def on_body(chunk: bytes, offset: int, **_kwargs: Any) -> None: on_headers=on_headers, on_body=on_body, ).finished_future.result() - return bytes(body) + return body def delete(self, key: str) -> None: self._client.make_request( @@ -155,11 +170,11 @@ def refit_env_int(name: str, *, default: int, min_value: int = 1) -> int: return value -def sparse_payload_checksum(body: bytes) -> str: +def sparse_payload_checksum(body: bytes | bytearray) -> str: return hashlib.blake2b(body, digest_size=16).hexdigest() -def decode_sparse_payload(body: bytes, checksum: str) -> bytes: +def decode_sparse_payload(body: bytes | bytearray, checksum: str) -> bytes: actual = sparse_payload_checksum(body) if actual != checksum: raise ValueError( @@ -197,6 +212,29 @@ def iter_sparse_weight_chunks( yield chunk, export_pull_s +def sparse_name_shard(name: str, shard_count: int) -> int: + return ( + int.from_bytes(hashlib.blake2b(name.encode(), digest_size=8).digest()) + % shard_count + ) + + +def vllm_refit_endpoints(base_urls: Sequence[str], path: str) -> list[str]: + return list( + dict.fromkeys( + f"{url.strip().rstrip('/')}{path}" for url in base_urls if url.strip() + ) + ) + + +def _partition_sparse_weights_by_name( + tensors: Iterable[NamedTensor], shard_rank: int, shard_count: int +) -> Iterator[NamedTensor]: + for name, tensor in tensors: + if sparse_name_shard(name, shard_count) == shard_rank: + yield name, tensor + + def refit_http_session() -> requests.Session: session = getattr(_CONTROL_SESSION_LOCAL, "session", None) if session is None: @@ -240,7 +278,7 @@ def sparse_export_chunk_size( ) -> int: requested = refit_env_int( f"NRL_REFIT_{transport.upper()}_EXPORT_CHUNK_BYTES", - default=(1024 if transport == "zmq" else 256) * 1024**2, + default=256 * 1024**2, min_value=1, ) if torch.cuda.is_available(): @@ -260,7 +298,10 @@ def init_sparse_delta_baseline_from_iterator( shard_rank: int, shard_count: int, transport: str, + partition: SparsePartitionMode = "chunks", ) -> None: + if partition == "names": + iterator = _partition_sparse_weights_by_name(iterator, shard_rank, shard_count) start_s = time.perf_counter() export_chunk_size = sparse_export_chunk_size(delta_tracker, transport) @@ -271,7 +312,7 @@ def init_sparse_delta_baseline_from_iterator( ): chunk_count = chunk_index + 1 export_pull_s += pull_s - if chunk_index % shard_count != shard_rank: + if partition == "chunks" and chunk_index % shard_count != shard_rank: continue started = time.perf_counter() delta_tracker.snapshot_baseline(chunk) @@ -294,29 +335,44 @@ def stream_sparse_delta_payloads( transfer_workers: int, shard_rank: int, shard_count: int, + partition: SparsePartitionMode = "chunks", ) -> SparseDeltaStreamResult: + if partition == "names": + iterator = _partition_sparse_weights_by_name(iterator, shard_rank, shard_count) prefix = transport.upper() encode_workers = refit_env_int( f"NRL_REFIT_{prefix}_ENCODE_WORKERS", default=max(2, min(8, os.cpu_count() or 8)), ) encode_executor = _executor(f"refit-{transport}-encode", encode_workers) + serialize_workers = min(4, encode_workers) + serialize_executor = _executor(f"refit-{transport}-serialize", serialize_workers) transfer_executor = _executor(f"refit-{transport}-transfer", transfer_workers) export_chunk_size = sparse_export_chunk_size(delta_tracker, transport) def encode_chunk( chunk: TensorBatch, - ) -> tuple[bytes | None, dict[str, float], int, int]: + ) -> tuple[TensorPayload | None, float, int, int, int]: + dense_bytes = sum(tensor.numel() * tensor.element_size() for _, tensor in chunk) started = time.perf_counter() payload, changed_elements, total_elements = ( delta_tracker.prepare_sparse_delta_payload(chunk) ) encode_s = time.perf_counter() - started - if not payload[2]: - return None, {"encode_s": encode_s}, changed_elements, total_elements + return ( + payload if payload[2] else None, + encode_s, + changed_elements, + total_elements, + dense_bytes, + ) + + def serialize_payloads( + payloads: tuple[TensorPayload, ...], encode_s: float + ) -> tuple[bytes, dict[str, float]]: started = time.perf_counter() buffer = io.BytesIO() - torch.save(payload, buffer) + torch.save(merge_sparse_payloads(payloads), buffer) raw_body = buffer.getvalue() serialize_s = time.perf_counter() - started started = time.perf_counter() @@ -329,8 +385,6 @@ def encode_chunk( "serialize_s": serialize_s, "compress_s": compress_s, }, - changed_elements, - total_elements, ) def transfer_payload( @@ -354,10 +408,14 @@ def transfer_payload( } chunk_count = 0 export_pull_s = 0.0 - encode_inflight: dict[Any, int] = {} + encode_inflight: set[Any] = set() + serialize_inflight: dict[Any, int] = {} transfer_inflight: set[Any] = set() worker_errors: list[Exception] = [] max_encode_inflight = encode_workers * 2 + max_serialize_inflight = serialize_workers * 2 + max_transfer_inflight = transfer_workers * 2 + bucket = _SparsePayloadBucket([]) def collect_transfers(*, block: bool) -> None: if not transfer_inflight: @@ -382,29 +440,70 @@ def collect_transfers(*, block: bool) -> None: receiver_timing, [result["receiver"]], maximum=False ) - def drain_encodes() -> None: - completed, _ = wait(encode_inflight, return_when=FIRST_COMPLETED) + def collect_serialized(*, block: bool) -> None: + if not serialize_inflight: + return + if block: + completed, _ = wait(serialize_inflight, return_when=FIRST_COMPLETED) + else: + completed = {future for future in serialize_inflight if future.done()} for future in completed: - payload_index = encode_inflight.pop(future) + index = serialize_inflight.pop(future) try: encoded = future.result() except Exception as error: worker_errors.append(error) continue - body, encode_timing, changed_elements, total_elements = encoded - counts["changed_elements"] += changed_elements - counts["total_elements"] += total_elements - if body is not None: - transfer_inflight.add( - transfer_executor.submit( - transfer_payload, - (body, encode_timing), - payload_index, - ) - ) + while len(transfer_inflight) >= max_transfer_inflight: + collect_transfers(block=True) + transfer_inflight.add( + transfer_executor.submit(transfer_payload, encoded, index) + ) collect_transfers(block=False) - payload_index = 0 + def submit_bucket() -> None: + if not bucket.payloads: + return + while len(serialize_inflight) >= max_serialize_inflight: + collect_serialized(block=True) + serialize_inflight[ + serialize_executor.submit( + serialize_payloads, tuple(bucket.payloads), bucket.encode_s + ) + ] = bucket.next_index + bucket.next_index += 1 + bucket.payloads.clear() + bucket.dense_bytes = 0 + bucket.encode_s = 0.0 + collect_serialized(block=False) + + def consume_encoded(encoded: Any) -> None: + payload, encode_s, changed_elements, total_elements, dense_bytes = encoded + counts["changed_elements"] += changed_elements + counts["total_elements"] += total_elements + if payload is None: + return + if ( + bucket.payloads + and bucket.dense_bytes + dense_bytes + > delta_tracker.sparse_bucket_size_bytes + ): + submit_bucket() + bucket.payloads.append(payload) + bucket.dense_bytes += dense_bytes + bucket.encode_s += encode_s + if bucket.dense_bytes >= delta_tracker.sparse_bucket_size_bytes: + submit_bucket() + + def drain_encodes() -> None: + completed, _ = wait(encode_inflight, return_when=FIRST_COMPLETED) + for future in completed: + encode_inflight.remove(future) + try: + consume_encoded(future.result()) + except Exception as error: + worker_errors.append(error) + stream_start = time.perf_counter() try: for chunk_index, (chunk, pull_s) in enumerate( @@ -412,29 +511,30 @@ def drain_encodes() -> None: ): chunk_count = chunk_index + 1 export_pull_s += pull_s - if chunk_index % shard_count != shard_rank: + if partition == "chunks" and chunk_index % shard_count != shard_rank: continue - if len(encode_inflight) >= max_encode_inflight: + while len(encode_inflight) >= max_encode_inflight: drain_encodes() - encode_inflight[encode_executor.submit(encode_chunk, chunk)] = payload_index - payload_index += 1 + encode_inflight.add(encode_executor.submit(encode_chunk, chunk)) while encode_inflight: drain_encodes() + submit_bucket() + while serialize_inflight: + collect_serialized(block=True) while transfer_inflight: collect_transfers(block=True) if worker_errors: raise worker_errors[0] except Exception: - for future in (*encode_inflight, *transfer_inflight): - future.cancel() - if encode_inflight: - wait(encode_inflight) - if transfer_inflight: - wait(transfer_inflight) + for futures in (encode_inflight, serialize_inflight, transfer_inflight): + for future in futures: + future.cancel() + if futures: + wait(futures) raise - timing = { + report = { "total_s": time.perf_counter() - stream_start, "export_pull_s": export_pull_s, **timing, @@ -445,16 +545,17 @@ def drain_encodes() -> None: "export_chunk_mb": export_chunk_size / 1e6, "shard_rank": shard_rank, "shard_count": shard_count, + "partition": partition, "changed_elements": counts["changed_elements"], "total_elements": counts["total_elements"], "changed_pct": 100.0 * counts["changed_elements"] / max(counts["total_elements"], 1), } - timing.update(receiver_timing) + report.update(receiver_timing) print( f"REFIT_{prefix}_TIMING " - + " ".join(f"{key}={value}" for key, value in timing.items()), + + " ".join(f"{key}={value}" for key, value in report.items()), flush=True, ) return { @@ -474,6 +575,7 @@ def stream_sparse_delta_payloads_via_s3_manifest( timeout_s: float, shard_rank: int, shard_count: int, + partition: SparsePartitionMode = "chunks", ) -> SparseDeltaStreamResult: urls = [url.strip().rstrip("/") for url in refit_targets if url.strip()] if not urls: @@ -485,7 +587,7 @@ def stream_sparse_delta_payloads_via_s3_manifest( bucket, os.getenv("NRL_REFIT_S3_REGION", "us-east-1").strip() or "us-east-1", ) - endpoint_urls = [f"{url}{G_VLLM_REFIT_S3_MANIFEST_PATH}" for url in urls] + endpoint_urls = vllm_refit_endpoints(urls, G_VLLM_REFIT_S3_MANIFEST_PATH) object_prefix = os.getenv("NRL_REFIT_S3_PREFIX", "nemo-rl-refit").strip("/") run_prefix = ( f"{object_prefix}/{transfer_id}/{shard_rank:06d}" @@ -533,12 +635,13 @@ def send_payload(body: bytes, payload_index: int) -> dict[str, Any]: ), shard_rank=shard_rank, shard_count=shard_count, + partition=partition, ) def post_vllm_refit_endpoints( endpoint_urls: Sequence[str], - body: Mapping[str, str] | bytes, + body: Mapping[str, Any] | bytes, *, api_key: str | None, timeout_s: float, @@ -566,7 +669,6 @@ def post(url: str) -> dict[str, Any]: pool = executor or _executor("refit-fanout", len(endpoint_urls)) futures = [pool.submit(post, url) for url in endpoint_urls] - wait(futures) return [future.result() for future in futures] @@ -576,18 +678,33 @@ def flush_vllm_refit_urls( api_key_env_var: str | None, timeout_s: float, ) -> list[dict[str, Any]]: - endpoint_urls = [ - f"{url}{G_VLLM_REFIT_FLUSH_PATH}" - for url in (url.strip().rstrip("/") for url in base_urls if url.strip()) - ] return post_vllm_refit_endpoints( - endpoint_urls, + vllm_refit_endpoints(base_urls, G_VLLM_REFIT_FLUSH_PATH), {}, api_key=vllm_refit_api_key(api_key_env_var), timeout_s=timeout_s, ) +def prepare_vllm_sparse_refit_urls( + base_urls: Sequence[str], + state_dict_info: Mapping[str, tuple[tuple[int, ...], torch.dtype]], + *, + api_key_env_var: str | None, + timeout_s: float, +) -> list[dict[str, Any]]: + tensors = { + name: [list(shape), str(dtype).removeprefix("torch.")] + for name, (shape, dtype) in state_dict_info.items() + } + return post_vllm_refit_endpoints( + vllm_refit_endpoints(base_urls, G_VLLM_REFIT_PREPARE_PATH), + {"tensors": tensors}, + api_key=vllm_refit_api_key(api_key_env_var), + timeout_s=timeout_s, + ) + + def download_s3_refit_payload( manifest: Mapping[str, Any], ) -> bytes: diff --git a/nemo_rl/utils/weight_transfer_sparse_codec.py b/nemo_rl/utils/weight_transfer_sparse_codec.py index d3b458dec37..667803452ed 100644 --- a/nemo_rl/utils/weight_transfer_sparse_codec.py +++ b/nemo_rl/utils/weight_transfer_sparse_codec.py @@ -17,6 +17,8 @@ import threading from collections.abc import Iterable, Mapping from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from math import prod from typing import Any, Literal import numpy as np @@ -28,6 +30,12 @@ SparseInfo = tuple[str, torch.Tensor, torch.Tensor, torch.Tensor, SparseOperation] TensorPayload = tuple[torch.Tensor, tuple[torch.Tensor, ...], list[dict[str, Any]]] PreparedTensorPayload = tuple[TensorPayload, int, int] +DecodedSparseItem = tuple[dict[str, Any], torch.Tensor, torch.Tensor] +DecodedSparsePayload = tuple[ + tuple[torch.Tensor, torch.Tensor], + tuple[torch.Tensor, ...], + list[dict[str, Any]], +] _INTEGER_DTYPE_BY_SIZE = { 1: torch.uint8, @@ -35,19 +43,103 @@ 4: torch.int32, 8: torch.int64, } -_DTYPE_BY_NAME = { - "bfloat16": torch.bfloat16, - "float16": torch.float16, - "float32": torch.float32, - "float64": torch.float64, - "float8_e4m3fn": torch.float8_e4m3fn, - "float8_e5m2": torch.float8_e5m2, - "int8": torch.int8, - "int16": torch.int16, - "int32": torch.int32, - "int64": torch.int64, - "uint8": torch.uint8, -} + + +class _TensorPayloadBuilder: + def __init__(self) -> None: + self.locations: list[torch.Tensor] = [] + self.value_parts: list[list[torch.Tensor]] = [] + self.value_group_by_dtype: dict[torch.dtype, int] = {} + self.value_offsets: list[int] = [] + self.metadata: list[dict[str, Any]] = [] + self.index_offset = 0 + + def add_locations(self, locations: torch.Tensor) -> tuple[int, int]: + start = self.index_offset + if locations.numel(): + self.locations.append(locations) + self.index_offset += locations.numel() + return start, self.index_offset + + def add_values(self, values: torch.Tensor) -> tuple[int, int, int]: + group = self.value_group_by_dtype.get(values.dtype) + if group is None: + group = len(self.value_parts) + self.value_group_by_dtype[values.dtype] = group + self.value_parts.append([]) + self.value_offsets.append(0) + start = self.value_offsets[group] + self.value_parts[group].append(values) + self.value_offsets[group] += values.numel() + return group, start, self.value_offsets[group] + + def finish(self) -> TensorPayload: + indices = ( + torch.cat(self.locations) + if self.locations + else torch.empty(0, dtype=torch.uint8) + ) + values = tuple( + torch.cat(parts) if len(parts) > 1 else parts[0] + for parts in self.value_parts + ) + return indices, values, self.metadata + + +@dataclass(frozen=True) +class SparseShardProjection: + """Map one local tensor shard into a canonical HF tensor.""" + + name: str + global_shape: tuple[int, ...] + offsets: tuple[int, ...] + + def map_locations( + self, locations: torch.Tensor, local_shape: tuple[int, ...] + ) -> torch.Tensor: + if len(local_shape) != len(self.global_shape) or len(local_shape) != len( + self.offsets + ): + raise ValueError(f"Sparse shard {self.name!r} has inconsistent ranks.") + if any( + offset < 0 or offset + local > global_size + for local, global_size, offset in zip( + local_shape, self.global_shape, self.offsets, strict=True + ) + ): + raise ValueError(f"Sparse shard {self.name!r} exceeds its global shape.") + if local_shape == self.global_shape and not any(self.offsets): + return locations + + mapped = torch.zeros_like(locations) + for dim, (local_size, global_size, offset) in enumerate( + zip(local_shape, self.global_shape, self.offsets, strict=True) + ): + local_stride = prod(local_shape[dim + 1 :]) + global_stride = prod(self.global_shape[dim + 1 :]) + coordinate = torch.div( + locations, local_stride, rounding_mode="floor" + ).remainder(local_size) + mapped.add_((coordinate + offset) * global_stride) + return mapped + + +@dataclass(frozen=True) +class SparseSourceRoute: + """Canonical source view consumed by one vLLM worker.""" + + offset: int + strides: tuple[int, ...] + shape: tuple[int, ...] + linear: bool + + +@dataclass(frozen=True) +class SparseSourcePlan: + """Source views needed by one worker for a canonical HF tensor.""" + + routes: tuple[SparseSourceRoute, ...] = () + identity: bool = False def integer_dtype_for_element_size(element_size: int) -> torch.dtype: @@ -58,10 +150,11 @@ def integer_dtype_for_element_size(element_size: int) -> torch.dtype: def dtype_from_name(name: str) -> torch.dtype: - try: - return _DTYPE_BY_NAME[name] - except KeyError as error: - raise ValueError(f"Unsupported sparse-refit tensor dtype {name!r}.") from error + dtype = getattr(torch, name, None) + if not isinstance(dtype, torch.dtype): + raise ValueError(f"Unsupported sparse-refit tensor dtype {name!r}.") + integer_dtype_for_element_size(dtype.itemsize) + return dtype def sparse_operation(value: object) -> SparseOperation: @@ -70,13 +163,6 @@ def sparse_operation(value: object) -> SparseOperation: raise ValueError(f"Unsupported sparse-refit operation {value!r}.") -def _dtype_name(dtype: torch.dtype) -> str: - name = str(dtype).removeprefix("torch.") - if name not in _DTYPE_BY_NAME: - raise ValueError(f"Unsupported sparse-refit tensor dtype {dtype}.") - return name - - def _integer_view(tensor: torch.Tensor) -> torch.Tensor: return tensor.contiguous().view( integer_dtype_for_element_size(tensor.element_size()) @@ -94,16 +180,11 @@ def _bytewise_diff_mask(current: torch.Tensor, baseline: torch.Tensor) -> torch. def encode_sparse_infos( infos: Iterable[SparseInfo], ) -> TensorPayload: - packed_locations = [] - value_parts: list[list[torch.Tensor]] = [] - value_group_by_dtype: dict[torch.dtype, int] = {} - value_offsets: list[int] = [] - metadata: list[dict[str, Any]] = [] - index_offset = 0 + payload = _TensorPayloadBuilder() for name, tensor, raw_locations, raw_values, operation in infos: count = int(raw_values.numel()) if count == 1 or int(raw_locations[-1] - raw_locations[0] + 1) == count: - index_count = 0 + index_start = index_end = payload.index_offset location_metadata = { "index_encoding": "range", "range_start": int(raw_locations[0]), @@ -111,41 +192,46 @@ def encode_sparse_infos( else: location_tensor = _encode_explicit_locations(raw_locations) location_metadata = {"index_encoding": "deltas"} - packed_locations.append(location_tensor) - index_count = int(location_tensor.numel()) - value_group = value_group_by_dtype.get(raw_values.dtype) - if value_group is None: - value_group = len(value_parts) - value_group_by_dtype[raw_values.dtype] = value_group - value_parts.append([]) - value_offsets.append(0) - value_start = value_offsets[value_group] - value_parts[value_group].append(raw_values) - value_offsets[value_group] += count - metadata.append( + index_start, index_end = payload.add_locations(location_tensor) + value_group, value_start, value_end = payload.add_values(raw_values) + payload.metadata.append( { "name": name, "shape": tuple(int(dim) for dim in tensor.shape), - "dtype": _dtype_name(tensor.dtype), + "dtype": str(tensor.dtype).removeprefix("torch."), "operation": operation, - "index_start": index_offset, - "index_end": index_offset + index_count, + "index_start": index_start, + "index_end": index_end, "value_group": value_group, "value_start": value_start, - "value_end": value_start + count, + "value_end": value_end, **location_metadata, } ) - index_offset += index_count - indices = ( - torch.cat(packed_locations) - if packed_locations - else torch.empty(0, dtype=torch.uint8) - ) - values = tuple( - torch.cat(parts) if len(parts) > 1 else parts[0] for parts in value_parts - ) - return indices, values, metadata + return payload.finish() + + +def merge_sparse_payloads(payloads: Iterable[TensorPayload]) -> TensorPayload: + """Combine encoded chunks without materializing dense source tensors.""" + payload = _TensorPayloadBuilder() + for locations, value_groups, items in payloads: + group_remap = {} + group_starts = {} + for old_group, values in enumerate(value_groups): + new_group, start, _ = payload.add_values(values) + group_remap[old_group] = new_group + group_starts[old_group] = start + index_offset, _ = payload.add_locations(locations) + for item in items: + merged = dict(item) + merged["index_start"] = int(item["index_start"]) + index_offset + merged["index_end"] = int(item["index_end"]) + index_offset + old_group = int(item["value_group"]) + merged["value_group"] = group_remap[old_group] + merged["value_start"] = int(item["value_start"]) + group_starts[old_group] + merged["value_end"] = int(item["value_end"]) + group_starts[old_group] + payload.metadata.append(merged) + return payload.finish() def sparse_locations_for_item( @@ -153,11 +239,14 @@ def sparse_locations_for_item( packed_locations: torch.Tensor, *, device: torch.device | int | str, + dtype: torch.dtype = torch.int64, ) -> torch.Tensor: + if dtype not in (torch.int32, torch.int64): + raise ValueError(f"Unsupported sparse location dtype {dtype}.") count = int(item["value_end"]) - int(item["value_start"]) if item["index_encoding"] == "range": start = int(item["range_start"]) - return torch.arange(start, start + count, device=device) + return torch.arange(start, start + count, dtype=dtype, device=device) index_start, index_end = int(item["index_start"]), int(item["index_end"]) raw = ( @@ -169,11 +258,168 @@ def sparse_locations_for_item( .tobytes() ) delta_dtype = {2: np.uint16, 4: np.uint32, 8: np.uint64}[len(raw) // count] - deltas = np.frombuffer(raw, dtype=delta_dtype).astype(np.int64, copy=False) - locations = np.cumsum(deltas + 1, dtype=np.int64) - 1 + location_dtype = np.int32 if dtype == torch.int32 else np.int64 + deltas = np.frombuffer(raw, dtype=delta_dtype).astype(location_dtype, copy=False) + locations = np.cumsum(deltas + 1, dtype=location_dtype) - 1 return torch.from_numpy(locations).to(device=device) +def _merge_tensor_parts(parts: list[torch.Tensor], dtype: torch.dtype) -> torch.Tensor: + if not parts: + return torch.empty(0, dtype=dtype) + return parts[0] if len(parts) == 1 else torch.cat(parts) + + +def decode_sparse_tensor_payload_for_staging( + payload: TensorPayload, +) -> DecodedSparsePayload: + """Flatten decoded locations so workers mmap only a few tensor storages.""" + packed_locations, value_groups, source_metadata = payload + location_parts: tuple[list[torch.Tensor], list[torch.Tensor]] = ([], []) + location_offsets = [0, 0] + metadata = [] + for item in source_metadata: + location_dtype = ( + torch.int32 + if prod(item["shape"]) <= torch.iinfo(torch.int32).max + else torch.int64 + ) + locations = sparse_locations_for_item( + item, packed_locations, device="cpu", dtype=location_dtype + ) + group = 0 if locations.dtype == torch.int32 else 1 + staged_item = dict(item) + staged_item["decoded_location_group"] = group + staged_item["decoded_location_start"] = location_offsets[group] + location_offsets[group] += locations.numel() + staged_item["decoded_location_end"] = location_offsets[group] + location_parts[group].append(locations) + metadata.append(staged_item) + locations = ( + _merge_tensor_parts(location_parts[0], torch.int32), + _merge_tensor_parts(location_parts[1], torch.int64), + ) + return locations, value_groups, metadata + + +def iter_decoded_sparse_payload( + payload: DecodedSparsePayload, +) -> Iterable[DecodedSparseItem]: + location_groups, value_groups, metadata = payload + for item in metadata: + location_start = int(item["decoded_location_start"]) + location_end = int(item["decoded_location_end"]) + value_start = int(item["value_start"]) + value_end = int(item["value_end"]) + yield ( + item, + location_groups[int(item["decoded_location_group"])][ + location_start:location_end + ], + value_groups[int(item["value_group"])][value_start:value_end], + ) + + +def map_sparse_locations( + locations: torch.Tensor, + source_offset: int, + source_strides: tuple[int, ...], + shape: tuple[int, ...], + linear: bool, + target_offset: int = 0, + target_strides: tuple[int, ...] | None = None, +) -> tuple[torch.Tensor, torch.Tensor]: + if linear: + end = source_offset + prod(shape) + keep = (locations >= source_offset) & (locations < end) + return locations + target_offset - source_offset, keep + mapped = torch.full_like(locations, target_offset) + reconstructed = torch.full_like(locations, source_offset) + relative = locations - source_offset + for size, source_stride, target_stride in zip( + shape, source_strides, target_strides or source_strides, strict=True + ): + coordinate = ( + torch.div(relative, source_stride, rounding_mode="floor").remainder(size) + if size > 1 + else torch.zeros_like(locations) + ) + reconstructed.add_(coordinate * source_stride) + mapped.add_(coordinate * target_stride) + return mapped, reconstructed == locations + + +def select_sparse_source_entries( + locations: torch.Tensor, + values: torch.Tensor, + plan: SparseSourcePlan, +) -> tuple[torch.Tensor, torch.Tensor]: + if plan.identity: + return locations, values + if not plan.routes: + return locations[:0], values[:0] + if all(route.linear for route in plan.routes): + ranges = sorted( + (route.offset, route.offset + prod(route.shape)) for route in plan.routes + ) + merged = [] + for start, end in ranges: + if merged and start <= merged[-1][1]: + merged[-1] = (merged[-1][0], max(merged[-1][1], end)) + else: + merged.append((start, end)) + bounds = torch.tensor( + [bound for interval in merged for bound in interval], + dtype=locations.dtype, + ) + offsets = torch.searchsorted(locations, bounds).tolist() + location_parts = [ + locations[start:end] for start, end in zip(offsets[::2], offsets[1::2]) + ] + value_parts = [ + values[start:end] for start, end in zip(offsets[::2], offsets[1::2]) + ] + if len(location_parts) == 1: + return location_parts[0], value_parts[0] + return torch.cat(location_parts), torch.cat(value_parts) + keep = torch.zeros(locations.shape, dtype=torch.bool) + for route in plan.routes: + _, route_keep = map_sparse_locations( + locations, route.offset, route.strides, route.shape, route.linear + ) + keep.logical_or_(route_keep) + return locations[keep], values[keep] + + +def partition_decoded_sparse_entries( + decoded: Iterable[DecodedSparseItem], + plans: Mapping[str, SparseSourcePlan], +) -> list[DecodedSparseItem]: + """Select the canonical entries consumed by one colocated worker.""" + selected_items = [] + for item, locations, values in decoded: + plan = plans[str(item["name"])] + selected_locations, selected_values = select_sparse_source_entries( + locations, values, plan + ) + if not selected_locations.numel(): + continue + selected_item = dict(item) + sample_location_tensor = torch.tensor( + item.get("verification_locations", ()), dtype=locations.dtype + ) + sample_value_tensor = torch.tensor( + item.get("verification_values", ()), dtype=values.dtype + ) + samples, sample_bits = select_sparse_source_entries( + sample_location_tensor, sample_value_tensor, plan + ) + selected_item["verification_locations"] = samples.tolist() + selected_item["verification_values"] = sample_bits.tolist() + selected_items.append((selected_item, selected_locations, selected_values)) + return selected_items + + def _encode_explicit_locations( locations: torch.Tensor, ) -> torch.Tensor: @@ -194,7 +440,12 @@ def _encode_explicit_locations( class DeltaCompressionTracker: """Source-side CPU or mmap baseline for sparse-delta refit.""" - def __init__(self, config: Mapping[str, Any]) -> None: + def __init__( + self, + config: Mapping[str, Any], + *, + projections: Mapping[str, SparseShardProjection] | None = None, + ) -> None: self.sparse_bucket_size_bytes = int(config["sparse_bucket_size_bytes"]) if self.sparse_bucket_size_bytes < 1: raise ValueError("delta_compression.sparse_bucket_size_bytes must be >= 1") @@ -206,6 +457,7 @@ def __init__(self, config: Mapping[str, Any]) -> None: raise ValueError("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD must be >= 0") self.baseline_in_memory = os.getenv("NRL_REFIT_BASELINE_IN_MEMORY") == "1" self.baseline_mmap_dir = os.getenv("NRL_REFIT_BASELINE_MMAP_DIR") + self.projections = dict(projections or {}) self.baseline: dict[str, torch.Tensor] = {} self._pending_updates: dict[str, tuple[torch.Tensor, torch.Tensor]] = {} self._pending_updates_lock = threading.Lock() @@ -223,35 +475,42 @@ def prepare_sparse_delta_payload( pending_updates = {} changed_elements = total_elements = 0 for name, tensor in tensors: - baseline = self.baseline.get(name) - if baseline is None: - raise RuntimeError(f"Sparse delta baseline is missing {name!r}.") - current = tensor.detach().cpu().contiguous() - current_bits = _integer_view(current).view(-1) + baseline, current, locations, current_values = self._find_changes( + name, tensor + ) baseline_bits = _integer_view(baseline).view(-1) total_elements += current.numel() - locations = ( - _bytewise_diff_mask(current, baseline).view(-1).nonzero().view(-1) - ) changed_elements += locations.numel() if locations.numel(): - current_values = current_bits[locations] values = ( current_values.bitwise_xor(baseline_bits[locations]) if self.encoding == "xor" else current_values ) + payload_name, payload_tensor, payload_locations = ( + name, + current, + locations, + ) + if projection := self.projections.get(name): + payload_name = projection.name + payload_locations = projection.map_locations( + locations, tuple(current.shape) + ) + payload_tensor = torch.empty( + projection.global_shape, dtype=current.dtype, device="meta" + ) sparse_infos.append( ( - name, - current, - locations, + payload_name, + payload_tensor, + payload_locations, values, self.encoding, ) ) if self.verification_samples: - verification_sources.append((locations, values)) + verification_sources.append((payload_locations, values)) pending_updates[name] = (locations, current_values) with self._pending_updates_lock: self._pending_updates.update(pending_updates) @@ -260,6 +519,36 @@ def prepare_sparse_delta_payload( self._add_verification_samples(payload[2], verification_sources) return payload, changed_elements, total_elements + def prepare_change_summary( + self, tensors: Iterable[NamedTensor] + ) -> tuple[set[str], int, int]: + """Scan local tensors without constructing a wire payload.""" + self._wait_for_baseline_commits() + changed_names = set() + pending_updates = {} + changed_elements = total_elements = 0 + for name, tensor in tensors: + _, current, locations, current_values = self._find_changes(name, tensor) + total_elements += current.numel() + changed_elements += locations.numel() + if locations.numel(): + changed_names.add(name) + pending_updates[name] = (locations, current_values) + with self._pending_updates_lock: + self._pending_updates.update(pending_updates) + return changed_names, changed_elements, total_elements + + def _find_changes( + self, name: str, tensor: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + baseline = self.baseline.get(name) + if baseline is None: + raise RuntimeError(f"Sparse delta baseline is missing {name!r}.") + current = tensor.detach().cpu().contiguous() + locations = _bytewise_diff_mask(current, baseline).view(-1).nonzero().view(-1) + current_values = _integer_view(current).view(-1)[locations] + return baseline, current, locations, current_values + def _add_verification_samples( self, metadata: list[dict[str, Any]], @@ -319,22 +608,10 @@ def _commit_baseline_updates( for name, (locations, values) in updates: target = _integer_view(self.baseline[name]).view(-1) count = locations.numel() - if count > 1: - first, last = int(locations[0]), int(locations[-1]) - span = last - first - if span % (count - 1) == 0: - step = span // (count - 1) - if step == 1 or all( - torch.equal( - locations[start:end], - first - + torch.arange(start, end, dtype=locations.dtype) * step, - ) - for start in range(0, count, 1 << 20) - for end in (min(start + (1 << 20), count),) - ): - target[first : last + 1 : step].copy_(values) - continue + first = int(locations[0]) + if int(locations[-1]) - first + 1 == count: + target[first : first + count].copy_(values) + continue target.index_copy_(0, locations, values) def _baseline( diff --git a/nemo_rl/utils/weight_transfer_zmq.py b/nemo_rl/utils/weight_transfer_zmq.py index c6b4f151330..624eba3a139 100644 --- a/nemo_rl/utils/weight_transfer_zmq.py +++ b/nemo_rl/utils/weight_transfer_zmq.py @@ -20,18 +20,21 @@ import uuid from collections.abc import Iterable, Mapping, Sequence from concurrent.futures import ThreadPoolExecutor +from contextlib import suppress from typing import Any import zmq from nemo_rl.utils.weight_transfer_remote_sparse import ( SparseDeltaStreamResult, + SparsePartitionMode, merge_vllm_refit_metrics, post_vllm_refit_endpoints, refit_env_int, sparse_payload_checksum, stream_sparse_delta_payloads, vllm_refit_api_key, + vllm_refit_endpoints, ) from nemo_rl.utils.weight_transfer_sparse_codec import ( DeltaCompressionTracker, @@ -161,11 +164,8 @@ def __init__( api_key_env_var: str | None, timeout_s: float, ) -> None: - self._refit_endpoints = tuple( - f"{url}{G_VLLM_REFIT_ZMQ_PAYLOAD_PATH}" - for url in dict.fromkeys( - url.strip().rstrip("/") for url in refit_urls if url.strip() - ) + self._refit_endpoints = vllm_refit_endpoints( + refit_urls, G_VLLM_REFIT_ZMQ_PAYLOAD_PATH ) if not self._refit_endpoints: raise ValueError("ZeroMQ sparse refit requires receiver HTTP URLs.") @@ -242,12 +242,10 @@ def _send_reply( kind: bytes, reply: Mapping[str, Any], ) -> None: - try: + with suppress(zmq.ZMQError): socket.send_multipart( [identity, kind, _json_bytes(reply)], flags=zmq.NOBLOCK ) - except zmq.ZMQError: - pass def _parse_data_message( self, @@ -299,28 +297,29 @@ def _run(self) -> None: self._ready.set() while not self._stop.is_set() or pending: - if not self._stop.is_set() and len(pending) < self._payload_workers: - if socket.poll(10, zmq.POLLIN): - frames = socket.recv_multipart() - identity = frames[0] if frames else b"" - try: - identity, key, body, metadata = self._parse_data_message( - frames - ) - future = payload_executor.submit( - self._fanout, - body, - metadata, - http_executor, - ) - pending[future] = (identity, key) - except Exception as exc: - self._send_reply( - socket, - identity, - _NACK, - {"ok": False, "error": str(exc)}, - ) + if ( + not self._stop.is_set() + and len(pending) < self._payload_workers + and socket.poll(10, zmq.POLLIN) + ): + frames = socket.recv_multipart() + identity = frames[0] if frames else b"" + try: + identity, key, body, metadata = self._parse_data_message(frames) + future = payload_executor.submit( + self._fanout, + body, + metadata, + http_executor, + ) + pending[future] = (identity, key) + except Exception as exc: + self._send_reply( + socket, + identity, + _NACK, + {"ok": False, "error": str(exc)}, + ) for future, (identity, key) in list(pending.items()): if not future.done(): @@ -361,6 +360,7 @@ def stream_sparse_delta_payloads_via_zmq( timeout_s: float, shard_rank: int, shard_count: int, + partition: SparsePartitionMode = "chunks", ) -> SparseDeltaStreamResult: addresses = [address.strip() for address in refit_targets if address.strip()] if not addresses: @@ -403,4 +403,5 @@ def send_payload(body: bytes, payload_id: int) -> dict[str, Any]: transfer_workers=refit_env_int("NRL_REFIT_ZMQ_SEND_WORKERS", default=4), shard_rank=shard_rank, shard_count=shard_count, + partition=partition, ) diff --git a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py index c16969ae47b..2ab0b40d73c 100644 --- a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py +++ b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py @@ -26,9 +26,15 @@ from nemo_rl.utils.weight_transfer_remote_sparse import ( flush_vllm_refit_urls, merge_vllm_refit_metrics, + prepare_vllm_sparse_refit_urls, ) from nemo_rl.weight_sync.interfaces import WeightSynchronizer +_REMOTE_SPARSE_TRANSPORTS = { + "vllm_s3_sparse": "s3", + "vllm_zmq_sparse": "zmq", +} + def validate_vllm_remote_sparse_refit( config: Any, @@ -36,9 +42,9 @@ def validate_vllm_remote_sparse_refit( colocated: bool, megatron_enabled: bool, ) -> str | None: - """Validate the optional transport without exposing its rules to GRPO.""" + """Validate the optional config and return its internal transport name.""" transport = config.get("refit_transport") - if transport not in (None, "vllm_s3_sparse", "vllm_zmq_sparse"): + if transport is not None and transport not in _REMOTE_SPARSE_TRANSPORTS: raise ValueError(f"Unsupported vLLM refit transport {transport!r}.") vllm_cfg = config["vllm_cfg"] if transport is not None and ( @@ -54,7 +60,7 @@ def validate_vllm_remote_sparse_refit( f"{transport} requires a non-colocated Megatron policy, BF16/FP16 " "vLLM, delta compression, and an unquantized rollout." ) - return transport + return None if transport is None else _REMOTE_SPARSE_TRANSPORTS[transport] class VllmRemoteSparseWeightSynchronizer(WeightSynchronizer): @@ -66,6 +72,7 @@ def __init__( transport: str, api_key_env_var: str | None = None, request_timeout_s: float = 600.0, + baseline_init_refs: list[Any] | None = None, ) -> None: self._policy = policy self._generation = generation @@ -74,7 +81,7 @@ def __init__( self._request_timeout_s = request_timeout_s self._refit_urls: list[str] = [] self._targets: list[str] = [] - self._baseline_init_refs: list[Any] = [] + self._baseline_init_refs = list(baseline_init_refs or ()) self._baseline_commit_refs: list[Any] = [] self._stale = True @@ -237,10 +244,31 @@ def _run_generation_workers(self, method_name: str, **kwargs: Any) -> list[Any]: ) ) - def init_communicator(self) -> None: - self._baseline_init_refs = self._run_policy_workers( - "init_remote_sparse_delta_baseline", transport=self._transport + @staticmethod + def start_baseline(policy: Any, transport: str) -> list[Any]: + workers = policy.worker_group + count = len(workers.workers) + return workers.run_all_workers_multiple_data( + "init_remote_sparse_delta_baseline", + common_kwargs={"transport": transport, "shard_count": count}, + shard_rank=list(range(count)), ) + + @staticmethod + def _merge_refit_info(parts: list[dict[str, Any]]) -> dict[str, Any]: + merged = {} + for part in parts: + for name, info in part.items(): + if name in merged and merged[name] != info: + raise ValueError(f"Conflicting sparse refit metadata for {name!r}.") + merged[name] = info + return merged + + def init_communicator(self) -> None: + if not self._baseline_init_refs: + self._baseline_init_refs = self.start_baseline( + self._policy, self._transport + ) self._refit_urls = [ url for url in self._run_generation_workers("report_refit_server_base_url") @@ -259,6 +287,16 @@ def init_communicator(self) -> None: raise ValueError( f"vLLM {self._transport} sparse refit endpoints are missing." ) + state_dict_info = self._merge_refit_info( + ray.get(list(self._baseline_init_refs)) + ) + self._baseline_init_refs.clear() + prepare_vllm_sparse_refit_urls( + self._refit_urls, + state_dict_info, + api_key_env_var=self._api_key_env_var, + timeout_s=self._request_timeout_s, + ) self._stale = False def shutdown(self) -> None: diff --git a/pyrefly.toml b/pyrefly.toml index 7e89fff230f..6458fc92148 100644 --- a/pyrefly.toml +++ b/pyrefly.toml @@ -171,6 +171,8 @@ project-includes = [ "nemo_rl/models/generation/vllm/reasoning_parsers/nano_v3_reasoning_parser.py", "nemo_rl/models/generation/vllm/utils.py", "nemo_rl/models/generation/vllm/vllm_backend.py", + "nemo_rl/models/generation/vllm/vllm_sparse_delta.py", + "nemo_rl/models/generation/vllm/vllm_sparse_refit.py", "nemo_rl/models/huggingface/__init__.py", "nemo_rl/models/megatron/__init__.py", "nemo_rl/models/megatron/draft/__init__.py", @@ -178,6 +180,7 @@ project-includes = [ "nemo_rl/models/policy/interfaces.py", "nemo_rl/models/policy/utils.py", "nemo_rl/models/policy/workers/__init__.py", + "nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py", "nemo_rl/models/policy/workers/patches.py", "nemo_rl/models/value/__init__.py", "nemo_rl/models/value/config.py", @@ -194,6 +197,9 @@ project-includes = [ "nemo_rl/utils/r3_trace.py", "nemo_rl/utils/timer.py", "nemo_rl/utils/venvs.py", + "nemo_rl/utils/weight_transfer_remote_sparse.py", + "nemo_rl/utils/weight_transfer_sparse_codec.py", + "nemo_rl/utils/weight_transfer_zmq.py", "nemo_rl/weight_sync/__init__.py", "nemo_rl/weight_sync/collective_weight_synchronizer.py", "nemo_rl/weight_sync/factory.py", @@ -201,19 +207,13 @@ project-includes = [ "nemo_rl/weight_sync/interfaces.py", "nemo_rl/weight_sync/ipc_weight_synchronizer.py", "nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py", - "nemo_rl/models/generation/vllm/vllm_sparse_delta.py", - "nemo_rl/models/generation/vllm/vllm_sparse_refit.py", - "nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py", - "nemo_rl/utils/weight_transfer_remote_sparse.py", - "nemo_rl/utils/weight_transfer_sparse_codec.py", - "nemo_rl/utils/weight_transfer_zmq.py", - "tools/refit_bandwidth_calculator.py", "tools/model_diagnostics/1.max_model_len_respected.py", "tools/model_diagnostics/2.long_generation_decode_vs_prefill.py", "tools/model_diagnostics/3.check_and_reinit_hf_model_embeddings_untrained.py", "tools/model_diagnostics/4.vllm_precision_compilation_test.py", "tools/model_diagnostics/5.prefix_caching_nan.py", "tools/model_diagnostics/6.vllm_routed_experts_completeness.py", + "tools/refit_bandwidth_calculator.py", "tools/x_token/__init__.py", "tools/x_token/reapply_exact_map.py", "tools/x_token/sort_and_cut_projection_matrix.py", diff --git a/tests/unit/models/generation/test_vllm_sparse_delta.py b/tests/unit/models/generation/test_vllm_sparse_delta.py index 757f16f9082..4be5fa61aae 100644 --- a/tests/unit/models/generation/test_vllm_sparse_delta.py +++ b/tests/unit/models/generation/test_vllm_sparse_delta.py @@ -15,6 +15,7 @@ import math from types import MethodType, SimpleNamespace from typing import Any +from unittest.mock import MagicMock import pytest import torch @@ -24,8 +25,13 @@ _SparseLoadTracer, ) from nemo_rl.utils.weight_transfer_sparse_codec import ( + SparseSourcePlan, + SparseSourceRoute, + decode_sparse_tensor_payload_for_staging, encode_sparse_infos, integer_dtype_for_element_size, + iter_decoded_sparse_payload, + partition_decoded_sparse_entries, ) @@ -83,39 +89,226 @@ def _bits(values: torch.Tensor) -> torch.Tensor: ) +def _decode_staged(payload: Any) -> list[Any]: + return list( + iter_decoded_sparse_payload(decode_sparse_tensor_payload_for_staging(payload)) + ) + + +def _apply_payload(applier: VllmSparseDeltaApplier, payload: Any) -> None: + decoded = _decode_staged(payload) + plans = applier.sparse_delta_source_plans([item for item, _, _ in decoded]) + applier._apply_decoded_sparse_weight_deltas( + partition_decoded_sparse_entries(decoded, plans) + ) + + +def test_sparse_plan_prewarm_uses_native_loader_without_applying_values() -> None: + target = torch.zeros(8) + applier = _applier(_NativeLoaderModel(identity=target)) + + applier.prewarm({"weight": ((8,), torch.float32)}) + + assert applier._plan_cache["weight"].identity + assert torch.equal(target, torch.zeros_like(target)) + + +def test_canonical_payload_is_partitioned_by_worker_source_plan() -> None: + tensor = torch.empty((8, 2), dtype=torch.float32) + locations = torch.tensor([1, 3, 8, 13]) + values = _bits(torch.tensor([1.0, 2.0, 3.0, 4.0])) + payload = encode_sparse_infos([("weight", tensor, locations, values, "overwrite")]) + payload[2][0].update( + verification_locations=[1, 13], + verification_values=[int(values[0]), int(values[3])], + ) + decoded = _decode_staged(payload) + + first = partition_decoded_sparse_entries( + decoded, + { + "weight": SparseSourcePlan( + routes=(SparseSourceRoute(0, (2, 1), (4, 2), True),) + ) + }, + ) + second = partition_decoded_sparse_entries( + decoded, + { + "weight": SparseSourcePlan( + routes=(SparseSourceRoute(8, (2, 1), (4, 2), True),) + ) + }, + ) + + first_item, first_locations, first_values = first[0] + second_item, second_locations, second_values = second[0] + assert first_locations.tolist() == [1, 3] + assert first_values.tolist() == values[:2].tolist() + assert first_item["verification_locations"] == [1] + assert second_locations.tolist() == [8, 13] + assert second_values.tolist() == values[2:].tolist() + assert second_item["verification_locations"] == [13] + + +def test_canonical_payload_partition_handles_strided_source_view() -> None: + values = torch.arange(8, dtype=torch.int32) + payload = encode_sparse_infos( + [ + ( + "weight", + torch.empty((8,), dtype=torch.float32), + torch.arange(8), + values, + "overwrite", + ) + ] + ) + partition = partition_decoded_sparse_entries( + _decode_staged(payload), + { + "weight": SparseSourcePlan( + routes=(SparseSourceRoute(1, (4, 1), (2, 2), False),) + ) + }, + ) + + _, locations, selected = partition[0] + assert locations.tolist() == [1, 2, 5, 6] + assert selected.tolist() == values[[1, 2, 5, 6]].tolist() + + +def test_canonical_payload_partitions_row_shards() -> None: + values = torch.arange(16, dtype=torch.int32) + payload = encode_sparse_infos( + [ + ( + "weight", + torch.empty((2, 8), dtype=torch.float32), + torch.arange(16), + values, + "overwrite", + ) + ] + ) + left = SparseSourcePlan(routes=(SparseSourceRoute(0, (8, 1), (2, 4), False),)) + right = SparseSourcePlan(routes=(SparseSourceRoute(4, (8, 1), (2, 4), False),)) + expected = {0: [0, 1, 2, 3, 8, 9, 10, 11], 1: [4, 5, 6, 7, 12, 13, 14, 15]} + for rank, plan in ((0, left), (1, right)): + _, locations, selected = partition_decoded_sparse_entries( + _decode_staged(payload), {"weight": plan} + )[0] + assert locations.tolist() == expected[rank] + assert selected.tolist() == values[expected[rank]].tolist() + + @pytest.mark.vllm -def test_serialized_sparse_payload_batch_preserves_order(tmp_path) -> None: +def test_backend_applies_decoded_sparse_payload_files() -> None: + from nemo_rl.models.generation.vllm.vllm_backend import ( + VllmInternalWorkerExtension, + ) + + ext = VllmInternalWorkerExtension.__new__(VllmInternalWorkerExtension) + applier = MagicMock() + applier.update_weights_from_decoded_sparse_payload.return_value = {"ok": True} + applier.update_weights_from_decoded_sparse_payload_files.return_value = {"ok": True} + ext._get_sparse_delta_applier = MagicMock(return_value=applier) + + ext.prepare_sparse_delta_refit_info({"weight": ((8,), torch.float32)}) + + assert ext.update_weights_from_decoded_sparse_payload(b"payload") == {"ok": True} + assert ext.update_weights_from_decoded_sparse_payload_files("first", "second") == { + "ok": True + } + applier.update_weights_from_decoded_sparse_payload.assert_called_once_with( + b"payload" + ) + applier.update_weights_from_decoded_sparse_payload_files.assert_called_once_with( + "first", "second" + ) + applier.prewarm.assert_called_once_with({"weight": ((8,), torch.float32)}) + + +@pytest.mark.vllm +def test_sparse_payload_batches_preserve_order(tmp_path) -> None: applier = VllmSparseDeltaApplier(SimpleNamespace(), torch.device("cpu")) - payloads = [ + decoded_paths = [tmp_path / f"decoded-{index}.pt" for index in range(3)] + decoded_payloads = [ ( - torch.tensor([index]), - (torch.tensor([float(index)]),), - [{"index": index}], + ( + torch.tensor([index], dtype=torch.int32), + torch.empty(0, dtype=torch.int64), + ), + (torch.tensor([index]),), + [ + { + "name": "weight", + "index": index, + "decoded_location_group": 0, + "decoded_location_start": 0, + "decoded_location_end": 1, + "value_group": 0, + "value_start": 0, + "value_end": 1, + } + ], ) for index in range(3) ] - paths = [tmp_path / f"{index}.pt" for index in range(3)] - for path, payload in zip(paths, payloads, strict=True): + for path, payload in zip(decoded_paths, decoded_payloads, strict=True): torch.save(payload, path) - applied: list[Any] = [] - compiled: list[Any] = [] - applier._compile_plans = compiled.append - applier._apply_sparse_weight_deltas = lambda tensors, metadata: applied.append( - (*tensors, metadata) - ) - - result = applier.update_weights_from_sparse_payload_files( - *(str(path) for path in paths) + decoded_applied: list[Any] = [] + applier.sparse_delta_source_plans = lambda _metadata: { + "weight": SparseSourcePlan(identity=True) + } + applier._apply_decoded_sparse_weight_deltas = decoded_applied.append + result = applier.update_weights_from_decoded_sparse_payload( + *(path.read_bytes() for path in decoded_paths) ) - applier.update_weights_from_serialized_sparse_payload( - *(path.read_bytes() for path in paths) + decoded_result = applier.update_weights_from_decoded_sparse_payload_files( + *(str(path) for path in reversed(decoded_paths)) ) - assert [[item["index"] for item in batch] for batch in compiled] == [[0, 1, 2]] * 2 - assert [item[2][0]["index"] for item in applied] == [0, 1, 2] * 2 + assert [payload[0][0]["index"] for payload in decoded_applied] == [ + 0, + 1, + 2, + 2, + 1, + 0, + ] assert result["receiver_deserialize_s"] >= 0.0 assert result["receiver_plan_s"] >= 0.0 assert result["receiver_sparse_apply_s"] >= 0.0 + assert decoded_result["receiver_deserialize_s"] >= 0.0 + assert decoded_result["receiver_partition_s"] >= 0.0 + + +@pytest.mark.vllm +def test_decoded_sparse_payload_converts_compact_locations_for_apply(tmp_path) -> None: + target = torch.zeros(8) + payload = encode_sparse_infos( + [ + ( + "weight", + target, + torch.tensor([1, 5]), + _bits(torch.tensor([2.0, 6.0])), + "overwrite", + ) + ] + ) + decoded = decode_sparse_tensor_payload_for_staging(payload) + assert decoded[0][0].dtype == torch.int32 + path = tmp_path / "decoded.pt" + torch.save(decoded, path) + + result = _applier( + _NativeLoaderModel(identity=target) + ).update_weights_from_decoded_sparse_payload_files(str(path)) + + assert torch.equal(target, torch.tensor([0.0, 2.0, 0.0, 0.0, 0.0, 6.0, 0.0, 0.0])) + assert result["receiver_partition_s"] >= 0.0 @pytest.mark.vllm @@ -174,7 +367,8 @@ def test_native_loaders_compile_sparse_placement() -> None: ) applier = _applier(_NativeLoaderModel(**targets)) - applier._apply_sparse_weight_deltas(payload[:2], payload[2]) + _apply_payload(applier, payload) + plans = applier.sparse_delta_source_plans(payload[2]) verification = applier.finish_sparse_delta_refit() assert torch.equal(targets["identity"], torch.tensor([0.0, 1.0, 2.0, 0.0])) @@ -184,7 +378,10 @@ def test_native_loaders_compile_sparse_placement() -> None: assert targets["w2"].view(-1)[[8, 11, 12, 15]].tolist() == [5, 5, 5, 5] assert targets["mamba"].view(-1)[[0, 3, 4, 11]].tolist() == [6, 6, 6, 6] assert torch.allclose(targets["a"], torch.tensor([-3.0, -2.0])) - assert verification["verification_candidates"] == 3 + assert plans["model.layers.0.self_attn.k_proj.weight"].routes == ( + SparseSourceRoute(4, (2, 1), (2, 2), True), + ) + assert verification["verification_candidates"] == 2 assert verification["verification_samples"] == 2 assert verification["verification_exact_mismatches"] == 0 @@ -244,9 +441,7 @@ def test_xor_applies_through_packed_native_loaders() -> None: ] ) - _applier(_NativeLoaderModel(**targets))._apply_sparse_weight_deltas( - payload[:2], payload[2] - ) + _apply_payload(_applier(_NativeLoaderModel(**targets)), payload) assert targets["qkv"].view(-1)[[8, 9, 11]].tolist() == [1.0, 1.0, 1.0] assert targets["merged"].view(-1)[[0, 1, 7, 8, 15]].tolist() == [2, 2, 2, 3, 3] @@ -361,7 +556,7 @@ def test_unknown_native_loader_fails_closed() -> None: ], ) with pytest.raises(RuntimeError, match=error): - _applier(model)._apply_sparse_weight_deltas(payload[:2], payload[2]) + _apply_payload(_applier(model), payload) @pytest.mark.vllm @@ -381,9 +576,7 @@ def test_unknown_sparse_operation_fails_closed() -> None: payload[2][0]["operation"] = "unknown" with pytest.raises(ValueError, match="Unsupported sparse-refit operation"): - _applier(_NativeLoaderModel(identity=target))._apply_sparse_weight_deltas( - payload[:2], payload[2] - ) + _apply_payload(_applier(_NativeLoaderModel(identity=target)), payload) @pytest.mark.vllm @@ -421,7 +614,7 @@ def test_sparse_delta_verification_compares_replacement( ) applier = _applier(_NativeLoaderModel(identity=target)) - applier._apply_sparse_weight_deltas(payload[:2], payload[2]) + _apply_payload(applier, payload) result = applier.finish_sparse_delta_refit() assert torch.equal(target, torch.tensor([1.0, initial + 4.0, 3.0, initial + 4.0])) @@ -469,8 +662,8 @@ def test_fp8_weight_and_scale_use_exact_bit_overwrite() -> None: ) applier = _applier(_NativeLoaderModel(identity=target, scale=scale)) - applier._apply_sparse_weight_deltas(payload[:2], payload[2]) - applier._apply_sparse_weight_deltas(payload[:2], payload[2]) + _apply_payload(applier, payload) + _apply_payload(applier, payload) result = applier.finish_sparse_delta_refit() assert target.view(torch.uint8).tolist() == [0x38, 0x41, 0x7F] @@ -495,14 +688,14 @@ def test_xor_applies_exact_bits_and_replay_reverts() -> None: ) applier = _applier(_NativeLoaderModel(identity=target)) - applier._apply_sparse_weight_deltas(payload[:2], payload[2]) + _apply_payload(applier, payload) result = applier.finish_sparse_delta_refit() assert torch.equal(_bits(target), _bits(current)) assert result["verification_exact_mismatches"] == 0 assert result["verification_mismatches"] == 0 - applier._apply_sparse_weight_deltas(payload[:2], payload[2]) + _apply_payload(applier, payload) assert torch.equal(_bits(target), _bits(baseline)) @@ -522,9 +715,7 @@ def test_overwrite_casts_absolute_source_values() -> None: ] ) - _applier(_NativeLoaderModel(identity=target))._apply_sparse_weight_deltas( - payload[:2], payload[2] - ) + _apply_payload(_applier(_NativeLoaderModel(identity=target)), payload) assert torch.equal(target, source.to(torch.float16)) @@ -558,9 +749,7 @@ def test_xor_rejects_non_bitwise_compatible_targets( ) with pytest.raises(RuntimeError, match=error): - _applier(_NativeLoaderModel(**targets))._apply_sparse_weight_deltas( - payload[:2], payload[2] - ) + _apply_payload(_applier(_NativeLoaderModel(**targets)), payload) @pytest.mark.vllm @@ -589,6 +778,4 @@ def load_weights(self, weights) -> None: ) with pytest.raises(RuntimeError, match="target mappings overlap"): - _applier(RepeatedCopyModel())._apply_sparse_weight_deltas( - payload[:2], payload[2] - ) + _apply_payload(_applier(RepeatedCopyModel()), payload) diff --git a/tests/unit/models/generation/test_vllm_sparse_refit.py b/tests/unit/models/generation/test_vllm_sparse_refit.py index a3b9e0f2fbd..259f859c9da 100644 --- a/tests/unit/models/generation/test_vllm_sparse_refit.py +++ b/tests/unit/models/generation/test_vllm_sparse_refit.py @@ -13,6 +13,7 @@ # limitations under the License. import asyncio +import io import threading import time from collections.abc import Iterator @@ -24,9 +25,11 @@ from unittest.mock import AsyncMock, MagicMock, call import pytest +import torch from nemo_rl.models.generation.vllm.vllm_sparse_refit import ( VllmSparseRefitReceiver, + _stage_sparse_payload, ) from nemo_rl.models.generation.vllm.vllm_worker_async import ( VllmAsyncGenerationWorkerImpl, @@ -34,8 +37,13 @@ from nemo_rl.utils.weight_transfer_remote_sparse import ( G_VLLM_REFIT_API_KEY_HEADER, G_VLLM_REFIT_FLUSH_PATH, + G_VLLM_REFIT_PREPARE_PATH, G_VLLM_REFIT_S3_MANIFEST_PATH, ) +from nemo_rl.utils.weight_transfer_sparse_codec import ( + encode_sparse_infos, + iter_decoded_sparse_payload, +) from nemo_rl.utils.weight_transfer_zmq import ( G_VLLM_REFIT_CHECKSUM_HEADER, G_VLLM_REFIT_PAYLOAD_HEADER, @@ -67,6 +75,42 @@ def _sparse_refit_receiver( yield receiver finally: receiver._refit_apply_executor.shutdown(wait=True) + receiver._refit_partition_executor.shutdown(wait=True) + + +def _serialized_sparse_payload() -> bytes: + values = torch.tensor([1, 2, 3, 4], dtype=torch.int32) + payload = encode_sparse_infos( + [ + ( + "weight", + torch.empty((8,), dtype=torch.float32), + torch.tensor([1, 3, 4, 7]), + values, + "overwrite", + ) + ] + ) + payload[2][0].update( + verification_locations=[1, 7], + verification_values=[1, 4], + ) + buffer = io.BytesIO() + torch.save(payload, buffer) + return buffer.getvalue() + + +def _stage_payloads( + receiver: VllmSparseRefitReceiver, + staging_dir: Path, + *payloads: bytes, +): + return tuple( + receiver._refit_partition_executor.submit( + _stage_sparse_payload, payload, str(staging_dir) + ) + for payload in payloads + ) def test_sparse_refit_queue_batches_payloads_in_fifo_order() -> None: @@ -100,6 +144,31 @@ def apply(payloads: tuple[bytes, ...]) -> dict[str, Any]: ) +def test_sparse_refit_queue_stages_payload_before_batch_is_full( + tmp_path: Path, +) -> None: + with _sparse_refit_receiver(batch_size=2) as receiver: + receiver._refit_workers_share_node = True + receiver._refit_batch_staging_dir = str(tmp_path) + + response = receiver._enqueue_sparse_payload_apply( + _serialized_sparse_payload(), ("transfer", 0, 0), "checksum" + ) + pending = receiver._refit_apply_pending_payloads[0] + assert isinstance(pending, Future) + staged = pending.result(timeout=1.0) + + assert response["ok"] + assert Path(staged.path).is_file() + assert receiver._refit_apply_futures == [] + receiver._flush_queued_sparse_payloads() + + assert not list(tmp_path.iterdir()) + assert receiver._worker.llm.collective_rpc.call_args_list[0].args[0] == ( + "update_weights_from_decoded_sparse_payload_files" + ) + + def test_sparse_refit_queue_deduplicates_transactional_payloads() -> None: key = ("transfer", 0, 1) with _sparse_refit_receiver() as receiver: @@ -189,37 +258,59 @@ def test_sparse_refit_queue_releases_condition_while_backpressured() -> None: assert pending_call.result(timeout=1.0)["ok"] -def test_sparse_refit_batch_uses_one_collective_rpc(tmp_path: Path) -> None: +def test_sparse_refit_batch_decodes_once_before_collective_apply( + tmp_path: Path, +) -> None: with _sparse_refit_receiver() as receiver: - staged_payloads: list[bytes] = [] - - def collective_rpc(method: str, args: tuple[str, ...]) -> list[dict[str, Any]]: - assert method == "update_weights_from_sparse_payload_files" - staged_payloads.extend(Path(path).read_bytes() for path in args) - return [{"ok": True, "receiver_total_s": 1.0}] + staged_locations: list[int] = [] + + def collective_rpc(method: str, args: tuple[Any, ...]) -> list[Any]: + assert method == "update_weights_from_decoded_sparse_payload_files" + for path in args: + staged_locations.extend( + int(location) + for _, locations, _ in iter_decoded_sparse_payload( + torch.load(path, weights_only=True) + ) + for location in locations + ) + return [ + {"ok": True, "receiver_total_s": 1.0}, + {"ok": True, "receiver_total_s": 1.0}, + {"ok": True, "receiver_total_s": 1.0}, + ] receiver._worker.llm = MagicMock( collective_rpc=MagicMock(side_effect=collective_rpc) ) - receiver._refit_workers_share_node = True - receiver._refit_batch_staging_dir = str(tmp_path) - response = receiver.update_weights_from_serialized_sparse_payloads( - (b"0", b"1", b"2") + response = receiver.update_weights_from_staged_sparse_payloads( + _stage_payloads(receiver, tmp_path, _serialized_sparse_payload()) + ) + receiver.update_weights_from_staged_sparse_payloads( + _stage_payloads(receiver, tmp_path, _serialized_sparse_payload()) ) - assert staged_payloads == [b"0", b"1", b"2"] + assert staged_locations == [1, 3, 4, 7] * 2 assert not list(tmp_path.iterdir()) - receiver._worker.llm.collective_rpc.assert_called_once() - assert response == {"ok": True, "receiver_total_s": 1.0, "payloads": 3} + assert [ + item.args[0] for item in receiver._worker.llm.collective_rpc.call_args_list + ] == [ + "update_weights_from_decoded_sparse_payload_files", + "update_weights_from_decoded_sparse_payload_files", + ] + assert response["payloads"] == 1 + assert response["receiver_worker_total_s"] == 1.0 + assert response["receiver_total_s"] >= 0.0 + assert receiver._refit_verification_candidates == 4 def test_sparse_refit_batch_drains_workers_before_error_cleanup(tmp_path: Path) -> None: with _sparse_refit_receiver() as receiver: staged_paths: tuple[str, ...] = () - def collective_rpc(method: str, args: tuple[str, ...]) -> list[Any]: + def collective_rpc(method: str, args: tuple[Any, ...]) -> list[Any]: nonlocal staged_paths - if method == "update_weights_from_sparse_payload_files": + if method == "update_weights_from_decoded_sparse_payload_files": staged_paths = args raise RuntimeError("apply failed") assert method == "synchronize_device" @@ -229,17 +320,16 @@ def collective_rpc(method: str, args: tuple[str, ...]) -> list[Any]: receiver._worker.llm = MagicMock( collective_rpc=MagicMock(side_effect=collective_rpc) ) - receiver._refit_workers_share_node = True - receiver._refit_batch_staging_dir = str(tmp_path) - with pytest.raises(RuntimeError, match="apply failed"): - receiver.update_weights_from_serialized_sparse_payloads((b"0", b"1")) + receiver.update_weights_from_staged_sparse_payloads( + _stage_payloads(receiver, tmp_path, _serialized_sparse_payload()) + ) assert [ entry.args[0] for entry in receiver._worker.llm.collective_rpc.call_args_list ] == [ - "update_weights_from_sparse_payload_files", + "update_weights_from_decoded_sparse_payload_files", "synchronize_device", ] assert not list(tmp_path.iterdir()) @@ -255,12 +345,15 @@ def test_sparse_refit_batch_uses_one_collective_rpc_across_nodes() -> None: ) response = receiver.update_weights_from_serialized_sparse_payloads( - (b"0", b"1", b"2") + (_serialized_sparse_payload(),) * 3 ) - receiver._worker.llm.collective_rpc.assert_called_once_with( - "update_weights_from_serialized_sparse_payload", - args=(b"0", b"1", b"2"), + rpc = receiver._worker.llm.collective_rpc.call_args + assert rpc.args[0] == "update_weights_from_decoded_sparse_payload" + assert len(rpc.kwargs["args"]) == 3 + assert all( + len(torch.load(io.BytesIO(payload), weights_only=True)[2]) == 1 + for payload in rpc.kwargs["args"] ) assert response == {"ok": True, "receiver_total_s": 1.0, "payloads": 3} @@ -270,28 +363,33 @@ async def test_async_sparse_refit_batch_bridges_to_async_collective( tmp_path: Path, ) -> None: with _sparse_refit_receiver(async_engine=True) as receiver: - staged_payloads: list[bytes] = [] + staged_locations: list[int] = [] class AsyncLlm: async def collective_rpc( - self, method: str, args: tuple[str, ...] - ) -> list[dict[str, Any]]: - assert method == "update_weights_from_sparse_payload_files" - staged_payloads.extend(Path(path).read_bytes() for path in args) + self, method: str, args: tuple[Any, ...] + ) -> list[Any]: + assert method == "update_weights_from_decoded_sparse_payload_files" + for path in args: + staged_locations.extend( + location + for _, locations, _ in iter_decoded_sparse_payload( + torch.load(path, weights_only=True) + ) + for location in locations.tolist() + ) return [{"ok": True, "receiver_total_s": 1.0}] receiver._worker.llm = AsyncLlm() receiver._refit_async_loop = asyncio.get_running_loop() - receiver._refit_workers_share_node = True - receiver._refit_batch_staging_dir = str(tmp_path) - response = await asyncio.to_thread( - receiver.update_weights_from_serialized_sparse_payloads, - (b"0", b"1"), + receiver.update_weights_from_staged_sparse_payloads, + _stage_payloads(receiver, tmp_path, _serialized_sparse_payload()), ) - assert staged_payloads == [b"0", b"1"] - assert response == {"ok": True, "receiver_total_s": 1.0, "payloads": 2} + assert staged_locations == [1, 3, 4, 7] + assert response["payloads"] == 1 + assert response["receiver_worker_total_s"] == 1.0 @pytest.mark.asyncio @@ -360,6 +458,7 @@ def test_sparse_refit_api_auth_dispatch_and_error_mapping(monkeypatch) -> None: receiver._flush_queued_sparse_payloads = MagicMock( return_value={"ok": True, "payloads": 2} ) + receiver._refit_collective_rpc = MagicMock(return_value=[]) app = FastAPI() receiver.setup_api_server(app) headers = {G_VLLM_REFIT_API_KEY_HEADER: "secret"} @@ -372,6 +471,11 @@ def test_sparse_refit_api_auth_dispatch_and_error_mapping(monkeypatch) -> None: headers=headers, ) flush_response = client.post(G_VLLM_REFIT_FLUSH_PATH, headers=headers) + prepare_response = client.post( + G_VLLM_REFIT_PREPARE_PATH, + json={"tensors": {"weight": [[2, 3], "bfloat16"]}}, + headers=headers, + ) zmq_response = client.post( G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, content=b"payload", @@ -381,11 +485,16 @@ def test_sparse_refit_api_auth_dispatch_and_error_mapping(monkeypatch) -> None: assert unauthorized.status_code == 403 assert s3_response.status_code == 200 assert flush_response.status_code == 200 + assert prepare_response.status_code == 200 assert zmq_response.status_code == 500 assert receiver._refit_async_loop is not None receiver._apply_s3_manifest_payload.assert_awaited_once_with({"key": "key"}) receiver._apply_zmq_payload.assert_awaited_once() receiver._flush_queued_sparse_payloads.assert_called_once_with() + receiver._refit_collective_rpc.assert_called_once_with( + "prepare_sparse_delta_refit_info", + ({"weight": ((2, 3), torch.bfloat16)},), + ) def test_sync_sparse_refit_server_shutdown_cleans_transport_resources( @@ -424,6 +533,8 @@ def run(self) -> None: } with _sparse_refit_receiver(config=config) as receiver: + receiver._worker.base_url = "http://10.0.0.2:8000/v1" + assert receiver.report_refit_server_base_url() == "http://10.0.0.2:8000" receiver.setup_api_server = MagicMock() receiver._setup_vllm_refit_server() assert configs[0].host == "0.0.0.0" diff --git a/tests/unit/models/policy/test_megatron_remote_sparse_refit.py b/tests/unit/models/policy/test_megatron_remote_sparse_refit.py index c6e33324214..f6fa839a3a0 100644 --- a/tests/unit/models/policy/test_megatron_remote_sparse_refit.py +++ b/tests/unit/models/policy/test_megatron_remote_sparse_refit.py @@ -12,23 +12,94 @@ # See the License for the specific language governing permissions and # limitations under the License. +from types import SimpleNamespace + import torch +from nemo_rl.models.policy.workers import megatron_remote_sparse_refit -def test_remote_sparse_stream_drains_cuda_before_return(monkeypatch): - from nemo_rl.models.policy.workers.megatron_remote_sparse_refit import ( - MegatronRemoteSparseRefit, +MegatronRemoteSparseRefit = megatron_remote_sparse_refit.MegatronRemoteSparseRefit + + +_DELTA_CONFIG = {"encoding": "overwrite", "sparse_bucket_size_bytes": 1024} +_XOR_CONFIG = {**_DELTA_CONFIG, "encoding": "xor"} + + +class _AutoMapping: + is_expert = False + is_adapter = False + is_grouped_export = False + ep_size = 1 + ep_rank = 0 + tp_size = 1 + tp_rank = 0 + + def __init__(self, hf_param, *, parallelism="column", permute_dims=None): + self.hf_param = hf_param + self.parallelism = parallelism + self.permute_dims = permute_dims + + def _detect_parallelism_type(self, _module): + return self.parallelism + + +class _ColumnMapping(_AutoMapping): + pass + + +class _DirectMapping(_AutoMapping): + pass + + +class _GatedMapping(_AutoMapping): + def __init__(self, *, gate, up): + super().__init__({"gate": gate, "up": up}) + + +class _RowMapping(_AutoMapping): + pass + + +class _ReplicatedMapping(_AutoMapping): + pass + + +def _install_mapping_types(monkeypatch, remote_refit_type): + monkeypatch.setattr( + remote_refit_type, + "_bridge_mapping_types", + staticmethod( + lambda: ( + _AutoMapping, + _ColumnMapping, + _DirectMapping, + _GatedMapping, + _ReplicatedMapping, + _RowMapping, + ) + ), + ) + monkeypatch.setattr( + remote_refit_type, + "_bridge_exports_are_identity", + lambda _self: True, ) + +def test_remote_sparse_stream_drains_cuda_before_return(monkeypatch): class Worker: + cfg = {} + fp8_cfg = None + model = object() + megatron_bridge = SimpleNamespace(get_conversion_tasks=lambda _models: []) + @staticmethod - def _iter_params_with_optional_kv_scales(): + def _iter_params_with_optional_kv_scales(*, conversion_tasks=None): + assert conversion_tasks == [] return iter(()) worker = Worker() - remote_refit = object.__new__(MegatronRemoteSparseRefit) - remote_refit._worker = worker - remote_refit._tracker = object() + remote_refit = MegatronRemoteSparseRefit(worker, _DELTA_CONFIG) result = {"payloads": 1, "changed_elements": 2, "total_elements": 3} events = [] @@ -58,3 +129,350 @@ def stream(*_args, **_kwargs): assert actual is result assert events == ["stream", "sync"] + + +def test_remote_sparse_stream_combines_local_and_misc_paths(monkeypatch): + class Worker: + @staticmethod + def _iter_params_with_optional_kv_scales(*, conversion_tasks=None): + assert conversion_tasks == [] + return iter(()) + + remote_refit = MegatronRemoteSparseRefit(Worker(), _DELTA_CONFIG) + remote_refit._local_tracker = object() + remote_refit._local_tensors = [("local", torch.ones(1))] + remote_refit._misc_conversion_tasks = [] + remote_refit._uses_policy_local_path = True + monkeypatch.setattr(remote_refit, "_changed_misc_tasks", lambda: ([], 5, 6)) + calls = [] + + def stream(*_args, **kwargs): + calls.append(kwargs) + if kwargs["partition"] == "names": + return {"payloads": 4, "changed_elements": 5, "total_elements": 6} + return {"payloads": 1, "changed_elements": 2, "total_elements": 3} + + monkeypatch.setattr( + "nemo_rl.models.policy.workers.megatron_remote_sparse_refit." + "stream_sparse_delta_payloads_via_zmq", + stream, + ) + monkeypatch.setattr(torch.cuda, "is_available", lambda: False) + + result = remote_refit.stream( + "zmq", + ["tcp://receiver:5555"], + transfer_id="transfer", + api_key_env_var=None, + timeout_s=1.0, + shard_rank=2, + shard_count=4, + ) + + assert result == {"payloads": 5, "changed_elements": 7, "total_elements": 9} + assert {call["transfer_id"] for call in calls} == { + "transfer-local", + "transfer-misc", + } + local = next(call for call in calls if call["transfer_id"].endswith("-local")) + assert local["partition"] == "none" + + +def test_remote_sparse_globalizes_expert_name(): + task = SimpleNamespace( + mapping=SimpleNamespace(is_expert=True, ep_size=4, ep_rank=2), + megatron_module=SimpleNamespace(config=SimpleNamespace(num_moe_experts=16)), + ) + + assert ( + MegatronRemoteSparseRefit._canonical_hf_name( + task, "model.layers.0.mlp.experts.1.up_proj.weight" + ) + == "model.layers.0.mlp.experts.9.up_proj.weight" + ) + assert ( + MegatronRemoteSparseRefit._canonical_hf_name( + task, "model.layers.0.mlp.experts.9.up_proj.weight" + ) + == "model.layers.0.mlp.experts.9.up_proj.weight" + ) + + +def test_remote_sparse_projects_bridge_affine_mappings(monkeypatch): + _install_mapping_types(monkeypatch, MegatronRemoteSparseRefit) + + def task(mapping, module_name, tensor, global_name): + module = type(module_name, (torch.nn.Module,), {})() + return SimpleNamespace( + mapping=mapping, + megatron_module=module, + param_weight=tensor, + global_param_name=global_name, + ) + + column = task( + _AutoMapping("backbone.layers.0.mixer.D"), + "TEColumnParallelLinear", + torch.arange(8).view(4, 2), + "decoder.layers.0.mixer.D", + ) + column.mapping.tp_size = 2 + column.mapping.tp_rank = 1 + row = task( + _AutoMapping("backbone.layers.0.mixer.o_proj.weight", parallelism="row"), + "TERowParallelLinear", + torch.arange(8).view(2, 4), + "decoder.layers.0.self_attention.linear_proj.weight", + ) + row.mapping.tp_size = 2 + row.mapping.tp_rank = 1 + replicated = task( + _AutoMapping("backbone.layers.0.norm.weight", parallelism="replicated"), + "TENorm", + torch.arange(4), + "decoder.layers.0.input_layernorm.weight", + ) + gated = task( + _GatedMapping( + gate="model.mlp.gate_proj.weight", + up="model.mlp.up_proj.weight", + ), + "TEColumnParallelLinear", + torch.arange(16).view(8, 2), + "mlp.linear_fc1.weight", + ) + + column_projection = MegatronRemoteSparseRefit._task_local_tensors( + column, megatron_remote_sparse_refit._COLUMN + )[0][2] + row_projection = MegatronRemoteSparseRefit._task_local_tensors( + row, megatron_remote_sparse_refit._ROW + )[0][2] + replicated_projection = MegatronRemoteSparseRefit._task_local_tensors( + replicated, megatron_remote_sparse_refit._REPLICATED + )[0][2] + gated_projections = MegatronRemoteSparseRefit._task_local_tensors( + gated, megatron_remote_sparse_refit._GATED + ) + + assert column_projection.name == "backbone.layers.0.mixer.D" + assert column_projection.global_shape == (8, 2) + assert column_projection.offsets == (4, 0) + assert row_projection.name == "backbone.layers.0.mixer.o_proj.weight" + assert row_projection.global_shape == (2, 8) + assert row_projection.offsets == (0, 4) + assert replicated_projection.name == "backbone.layers.0.norm.weight" + assert replicated_projection.global_shape == (4,) + assert replicated_projection.offsets == (0,) + assert [projection.name for _, _, projection in gated_projections] == [ + "model.mlp.gate_proj.weight", + "model.mlp.up_proj.weight", + ] + assert [tuple(tensor.shape) for _, tensor, _ in gated_projections] == [ + (4, 2), + (4, 2), + ] + + +def test_remote_sparse_uses_local_baseline_to_gate_transformed_tasks( + monkeypatch, +): + class _TransformedMapping(_AutoMapping): + pass + + tensor = torch.tensor([1.0, 2.0, 3.0]) + task = SimpleNamespace( + mapping=_TransformedMapping("model.q_proj.weight"), + megatron_module=torch.nn.Linear(3, 1), + param_weight=tensor, + global_param_name="decoder.linear_qkv.weight", + ) + + class Worker: + cfg = {} + fp8_cfg = None + model = object() + megatron_bridge = SimpleNamespace(get_conversion_tasks=lambda _models: [task]) + + _install_mapping_types(monkeypatch, MegatronRemoteSparseRefit) + monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") + remote_refit = MegatronRemoteSparseRefit(Worker(), _DELTA_CONFIG) + remote_refit._prepare_paths() + + assert remote_refit._misc_conversion_tasks == [task] + assert remote_refit._misc_local_tensors == [("0:decoder.linear_qkv.weight", tensor)] + assert remote_refit._change_tracker is not None + remote_refit._change_tracker.snapshot_baseline(remote_refit._misc_local_tensors) + + tensor[1] = 5 + changed_tasks, changed, total = remote_refit._changed_misc_tasks() + assert changed_tasks == [task] + assert (changed, total) == (1, 3) + + remote_refit._change_tracker.on_sync_succeeded() + assert remote_refit._changed_misc_tasks() == ([], 0, 3) + + +def test_remote_sparse_uses_xor_only_for_direct_bitwise_path(monkeypatch): + class _TransformedMapping(_AutoMapping): + pass + + direct = torch.tensor([[1.0, 2.0], [3.0, 4.0]]) + transformed = torch.tensor([1.0, 2.0]) + tasks = [ + SimpleNamespace( + mapping=_AutoMapping("model.q_proj.weight"), + megatron_module=torch.nn.Linear(2, 2), + param_weight=direct, + global_param_name="decoder.linear_qkv.weight", + ), + SimpleNamespace( + mapping=_TransformedMapping("backbone.layers.0.mixer.A_log"), + megatron_module=torch.nn.Linear(2, 1), + param_weight=transformed, + global_param_name="decoder.mixer.A_log", + ), + ] + + class Worker: + cfg = {} + fp8_cfg = None + model = object() + megatron_bridge = SimpleNamespace(get_conversion_tasks=lambda _models: tasks) + + _install_mapping_types(monkeypatch, MegatronRemoteSparseRefit) + monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") + remote_refit = MegatronRemoteSparseRefit(Worker(), _XOR_CONFIG) + remote_refit._prepare_paths() + + assert remote_refit._local_tracker is not None + assert remote_refit._local_tracker.encoding == "xor" + assert remote_refit._tracker.encoding == "overwrite" + + remote_refit._local_tracker.snapshot_baseline(remote_refit._local_tensors) + direct[0, 0] = 5 + direct_metadata = remote_refit._local_tracker.prepare_sparse_delta_payload( + remote_refit._local_tensors + )[0][2] + assert [item["operation"] for item in direct_metadata] == ["xor"] + + residual = [("backbone.layers.0.mixer.A_log", transformed)] + remote_refit._tracker.snapshot_baseline(residual) + transformed[0] = 3 + residual_metadata = remote_refit._tracker.prepare_sparse_delta_payload(residual)[0][ + 2 + ] + assert [item["operation"] for item in residual_metadata] == ["overwrite"] + assert remote_refit.refit_info() == { + "backbone.layers.0.mixer.A_log": ((2,), torch.float32), + "model.q_proj.weight": ((2, 2), torch.float32), + } + + +def test_remote_sparse_preserves_bridge_task_dependencies(monkeypatch): + grouped = SimpleNamespace(is_grouped_export=True, group_key="experts") + tasks = [ + SimpleNamespace(mapping=grouped), + SimpleNamespace(mapping=grouped), + SimpleNamespace( + mapping=SimpleNamespace(is_grouped_export=False, group_key="other") + ), + ] + remote_refit = MegatronRemoteSparseRefit(object(), _DELTA_CONFIG) + remote_refit._misc_conversion_tasks = tasks + remote_refit._filter_misc_tasks = True + monkeypatch.setattr(remote_refit, "_all_reduce_max", lambda _flags: [1, 0, 0]) + + assert remote_refit._changed_misc_tasks()[0] == tasks[:2] + + remote_refit._filter_misc_tasks = False + assert remote_refit._changed_misc_tasks()[0] == tasks + + +def test_remote_sparse_balances_tasks_across_equivalent_replicas(monkeypatch): + from megatron.core import parallel_state + + from nemo_rl.utils.weight_transfer_remote_sparse import sparse_name_shard + + ranks = {"dp": 0, "expert_dp": 0} + monkeypatch.setattr(torch.distributed, "is_initialized", lambda: True) + monkeypatch.setattr( + parallel_state, + "get_data_parallel_rank", + lambda *, with_context_parallel: ranks["dp"], + ) + monkeypatch.setattr( + parallel_state, + "get_data_parallel_world_size", + lambda *, with_context_parallel: 2, + ) + monkeypatch.setattr( + parallel_state, + "get_expert_data_parallel_rank", + lambda: ranks["expert_dp"], + ) + monkeypatch.setattr( + parallel_state, + "get_expert_data_parallel_world_size", + lambda: 4, + ) + + dense = SimpleNamespace( + global_param_name="decoder.layers.0.input_layernorm.weight", + mapping=SimpleNamespace(is_expert=False, tp_rank=0, tp_size=2), + ) + dense_owners = [] + for dp_rank in range(2): + ranks["dp"] = dp_rank + for tp_rank in range(2): + dense.mapping.tp_rank = tp_rank + if MegatronRemoteSparseRefit._owns_policy_local_task( + dense, replicated=True + ): + dense_owners.append(dp_rank * 2 + tp_rank) + assert dense_owners == [sparse_name_shard(dense.global_param_name, 4)] + + owner_counts = [0] * 4 + for expert_id in range(64): + name = f"decoder.layers.0.mlp.experts.local_experts.{expert_id}.weight" + expert = SimpleNamespace( + global_param_name=name, + mapping=SimpleNamespace(is_expert=True), + ) + owners = [] + for expert_dp_rank in range(4): + ranks["expert_dp"] = expert_dp_rank + if MegatronRemoteSparseRefit._owns_policy_local_task(expert): + owners.append(expert_dp_rank) + owner_counts[expert_dp_rank] += 1 + assert owners == [sparse_name_shard(name, 4)] + assert min(owner_counts) > 0 + + +def test_remote_sparse_fp8_policy_keeps_full_export_path(monkeypatch): + task = object() + + class Worker: + cfg = {} + fp8_cfg = {"fp8_param": True} + model = object() + megatron_bridge = SimpleNamespace(get_conversion_tasks=lambda _models: [task]) + + @staticmethod + def _iter_params_with_optional_kv_scales(*, conversion_tasks): + assert conversion_tasks == [task] + return iter(()) + + remote_refit = MegatronRemoteSparseRefit(Worker(), _DELTA_CONFIG) + snapshots = [] + monkeypatch.setattr( + "nemo_rl.models.policy.workers.megatron_remote_sparse_refit." + "init_sparse_delta_baseline_from_iterator", + lambda iterator, **_kwargs: snapshots.append(list(iterator)), + ) + + remote_refit.initialize_baseline(shard_rank=0, shard_count=1, transport="zmq") + + assert snapshots == [[]] + assert remote_refit._misc_conversion_tasks == [task] + assert not remote_refit._uses_policy_local_path diff --git a/tests/unit/utils/test_weight_transfer_remote_sparse.py b/tests/unit/utils/test_weight_transfer_remote_sparse.py index cac740aa5a2..f01f2da670a 100644 --- a/tests/unit/utils/test_weight_transfer_remote_sparse.py +++ b/tests/unit/utils/test_weight_transfer_remote_sparse.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import io import json import threading from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer @@ -22,9 +23,13 @@ import zstandard from nemo_rl.utils import weight_transfer_remote_sparse, weight_transfer_zmq -from nemo_rl.utils.weight_transfer_remote_sparse import download_s3_refit_payload +from nemo_rl.utils.weight_transfer_remote_sparse import ( + download_s3_refit_payload, + sparse_payload_checksum, +) from nemo_rl.utils.weight_transfer_sparse_codec import ( DeltaCompressionTracker, + SparseShardProjection, _bytewise_diff_mask, encode_sparse_infos, sparse_locations_for_item, @@ -36,7 +41,6 @@ G_VLLM_REFIT_TRANSFER_HEADER, ZmqSparseRefitClient, ZmqSparseRefitServer, - sparse_payload_checksum, ) @@ -45,7 +49,20 @@ class _SparsePipelineTracker: @staticmethod def prepare_sparse_delta_payload(chunk): - return (chunk, torch.ones(1), [1]), 1, 1 + count = sum(tensor.numel() for _, tensor in chunk) + payload = encode_sparse_infos( + ( + ( + name, + tensor, + torch.arange(tensor.numel()), + tensor.reshape(-1), + "overwrite", + ) + for name, tensor in chunk + ) + ) + return payload, count, count def _stream_sparse_test_payloads(tensors, send_payload): @@ -80,6 +97,22 @@ def test_delta_tracker_commits_only_successful_syncs(monkeypatch) -> None: assert not tracker.prepare_sparse_delta_payload([("weight", tensor)])[0][2] +def test_delta_tracker_change_summary_is_transactional(monkeypatch) -> None: + monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") + tracker = _delta_tracker() + first = torch.tensor([1.0, 2.0]) + second = torch.tensor([3.0, 4.0]) + tensors = [("first", first), ("second", second)] + tracker.snapshot_baseline(tensors) + second[0] = 5 + + assert tracker.prepare_change_summary(tensors) == ({"second"}, 1, 4) + tracker.on_sync_failed() + assert tracker.prepare_change_summary(tensors) == ({"second"}, 1, 4) + tracker.on_sync_succeeded() + assert tracker.prepare_change_summary(tensors) == (set(), 0, 4) + + def test_delta_tracker_emits_bounded_delta_samples(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") monkeypatch.setenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "2") @@ -142,6 +175,45 @@ def test_delta_tracker_xor_encodes_against_baseline(monkeypatch) -> None: assert torch.equal(tracker.baseline["weight"], tensor) +@pytest.mark.parametrize( + ("projection", "expected_locations"), + [ + (SparseShardProjection("hf.weight", (2, 4), (0, 2)), [2, 7]), + (SparseShardProjection("hf.weight", (4, 2), (2, 0)), [4, 7]), + ], +) +def test_delta_tracker_projects_local_shards_to_hf_locations( + monkeypatch, + projection: SparseShardProjection, + expected_locations: list[int], +) -> None: + monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") + monkeypatch.setenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "2") + tracker = DeltaCompressionTracker( + {"encoding": "overwrite", "sparse_bucket_size_bytes": 1024}, + projections={"local.weight": projection}, + ) + tensor = torch.zeros(2, 2) + tracker.snapshot_baseline([("local.weight", tensor)]) + tensor[0, 0] = 1 + tensor[1, 1] = 2 + + (locations, _, metadata), changed, total = tracker.prepare_sparse_delta_payload( + [("local.weight", tensor)] + ) + + assert (changed, total) == (2, 4) + assert metadata[0]["name"] == "hf.weight" + assert metadata[0]["shape"] == projection.global_shape + assert metadata[0]["verification_locations"] == expected_locations + assert ( + sparse_locations_for_item(metadata[0], locations, device="cpu").tolist() + == expected_locations + ) + tracker.on_sync_succeeded() + assert not tracker.prepare_sparse_delta_payload([("local.weight", tensor)])[0][2] + + def test_sparse_index_encoding_preserves_uint64_locations() -> None: locations = torch.tensor([0, 2**32 + 5]) packed, _, metadata = encode_sparse_infos( @@ -230,7 +302,7 @@ def test_s3_download_verifies_checksum(monkeypatch) -> None: monkeypatch.setattr( weight_transfer_remote_sparse, "_get_manifest_s3_store", - lambda *_args: SimpleNamespace(get=lambda _key: compressed), + lambda *_args: SimpleNamespace(get=lambda _key: bytearray(compressed)), ) manifest = { "bucket": "bucket", @@ -256,6 +328,34 @@ def test_refit_http_session_does_not_retry_application_errors() -> None: assert {502, 503, 504} <= set(retry.status_forcelist) +def test_prepare_sparse_refit_urls_serializes_metadata(monkeypatch) -> None: + posts = [] + monkeypatch.setenv("NRL_TEST_REFIT_KEY", "secret") + monkeypatch.setattr( + weight_transfer_remote_sparse, + "post_vllm_refit_endpoints", + lambda *args, **kwargs: posts.append((args, kwargs)) or [{"ok": True}], + ) + + result = weight_transfer_remote_sparse.prepare_vllm_sparse_refit_urls( + [" http://receiver/ "], + {"weight": ((2, 3), torch.bfloat16)}, + api_key_env_var="NRL_TEST_REFIT_KEY", + timeout_s=7.0, + ) + + assert result == [{"ok": True}] + assert posts == [ + ( + ( + ["http://receiver/nemo-rl/refit/prepare"], + {"tensors": {"weight": [[2, 3], "bfloat16"]}}, + ), + {"api_key": "secret", "timeout_s": 7.0}, + ) + ] + + def test_sparse_export_finishes_before_blocked_transfers(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_ZMQ_ENCODE_WORKERS", "1") monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "1") @@ -306,6 +406,53 @@ def fail_transfer(_body, _payload_index): assert exported == list(range(4)) +def test_sparse_stream_coalesces_export_chunks(monkeypatch) -> None: + monkeypatch.setenv("NRL_REFIT_ZMQ_ENCODE_WORKERS", "2") + monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "4") + tracker = _delta_tracker() + tracker.sparse_bucket_size_bytes = 8 + tensors = [(f"weight-{index}", torch.zeros(1)) for index in range(4)] + tracker.snapshot_baseline(tensors) + for _, tensor in tensors: + tensor.fill_(1) + payloads = {} + + def send(body, payload_index): + raw = zstandard.ZstdDecompressor().decompress(body) + payloads[payload_index] = torch.load( + io.BytesIO(raw), map_location="cpu", weights_only=True + ) + return {"receiver": {}} + + result = weight_transfer_remote_sparse.stream_sparse_delta_payloads( + tensors, + delta_tracker=tracker, + transport="zmq", + send_payload=send, + transfer_workers=1, + shard_rank=0, + shard_count=1, + partition="none", + ) + + assert result == {"payloads": 2, "changed_elements": 4, "total_elements": 4} + payload_names = [ + [item["name"] for item in payloads[index][2]] for index in sorted(payloads) + ] + assert all(len(names) == 2 for names in payload_names) + assert sorted(name for names in payload_names for name in names) == [ + f"weight-{index}" for index in range(4) + ] + for locations, value_groups, metadata in payloads.values(): + for item in metadata: + assert sparse_locations_for_item( + item, locations, device="cpu" + ).tolist() == [0] + assert value_groups[item["value_group"]][ + item["value_start"] : item["value_end"] + ].tolist() == [1065353216] + + def test_sparse_baseline_snapshots_only_owned_export_chunks( monkeypatch, capsys ) -> None: @@ -332,6 +479,60 @@ def snapshot_baseline(self, chunk) -> None: assert tracker.names == ["weight-1", "weight-3"] assert "chunks=4" in capsys.readouterr().out + tracker = Tracker() + weight_transfer_remote_sparse.init_sparse_delta_baseline_from_iterator( + [(f"weight-{index}", torch.tensor([float(index)])) for index in range(4)], + delta_tracker=tracker, + shard_rank=1, + shard_count=2, + transport="zmq", + partition="none", + ) + assert tracker.names == [f"weight-{index}" for index in range(4)] + + +def test_sparse_name_partition_is_stable_for_filtered_exports(monkeypatch) -> None: + monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "4") + + class Tracker: + sparse_bucket_size_bytes = 4 + + def __init__(self) -> None: + self.names = [] + + def snapshot_baseline(self, chunk) -> None: + self.names.extend(name for name, _tensor in chunk) + + tensors = [(f"weight-{index}", torch.tensor([float(index)])) for index in range(8)] + owners = [] + for rank in range(2): + tracker = Tracker() + weight_transfer_remote_sparse.init_sparse_delta_baseline_from_iterator( + tensors, + delta_tracker=tracker, + shard_rank=rank, + shard_count=2, + transport="zmq", + partition="names", + ) + owners.append(set(tracker.names)) + + assert owners[0].isdisjoint(owners[1]) + assert owners[0] | owners[1] == {name for name, _tensor in tensors} + + filtered = tensors[::2] + for rank in range(2): + tracker = Tracker() + weight_transfer_remote_sparse.init_sparse_delta_baseline_from_iterator( + filtered, + delta_tracker=tracker, + shard_rank=rank, + shard_count=2, + transport="zmq", + partition="names", + ) + assert set(tracker.names) == owners[rank] & {name for name, _tensor in filtered} + def test_s3_manifest_transport_validates_configuration(monkeypatch) -> None: kwargs = { 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 5107241abb3..04ffe5962c5 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 @@ -18,6 +18,7 @@ from nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer import ( VllmRemoteSparseWeightSynchronizer, + validate_vllm_remote_sparse_refit, ) @@ -36,7 +37,6 @@ def _remote_sparse_sync( policy.worker_group.run_all_workers_single_data.return_value = commit_refs generation = MagicMock() - generation.worker_group.workers = [object()] generation.worker_group.run_all_workers_single_data.side_effect = [ [MagicMock()], *([[MagicMock()]] if transport == "zmq" else []), @@ -46,15 +46,103 @@ def _remote_sparse_sync( get_results: list[object] = [["http://receiver"]] if transport == "zmq": get_results.append(["tcp://relay:19090"]) - get_results.extend([None, stream_result]) + get_results.extend([[{"weight": ((8,), "float32")}], stream_result]) mock_ray.get.side_effect = get_results sync = VllmRemoteSparseWeightSynchronizer(policy, generation, transport=transport) - sync.init_communicator() + with patch( + "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer." + "prepare_vllm_sparse_refit_urls" + ): + sync.init_communicator() return sync, policy, generation +def _valid_config() -> dict: + return { + "refit_transport": "vllm_s3_sparse", + "delta_compression": {"encoding": "overwrite"}, + "vllm_cfg": {"precision": "bfloat16", "kv_cache_dtype": "auto"}, + } + + +def test_validate_remote_sparse_refit_accepts_supported_scope(): + assert ( + validate_vllm_remote_sparse_refit( + _valid_config(), colocated=False, megatron_enabled=True + ) + == "s3" + ) + + +@pytest.mark.parametrize( + ("change", "kwargs"), + [ + ({"refit_transport": "unknown"}, {}), + ({}, {"colocated": True}), + ({}, {"megatron_enabled": False}), + ({"delta_compression": None}, {}), + ({"quant_cfg": "fp8"}, {}), + ({"vllm_cfg": {"precision": "fp8", "kv_cache_dtype": "auto"}}, {}), + ( + {"vllm_cfg": {"precision": "bfloat16", "kv_cache_dtype": "fp8_e4m3"}}, + {}, + ), + ], +) +def test_validate_remote_sparse_refit_rejects_unsupported_scope(change, kwargs): + config = _valid_config() + config.update(change) + arguments = {"colocated": False, "megatron_enabled": True} + arguments.update(kwargs) + + with pytest.raises(ValueError): + validate_vllm_remote_sparse_refit(config, **arguments) + + class TestVllmRemoteSparseWeightSynchronizer: + @patch( + "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer." + "prepare_vllm_sparse_refit_urls" + ) + @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") + def test_init_communicator_joins_prelaunched_baseline(self, mock_ray, prepare): + baseline_ref = MagicMock() + policy = MagicMock() + generation = MagicMock() + generation.worker_group.run_all_workers_single_data.return_value = [MagicMock()] + mock_ray.get.side_effect = [ + ["http://receiver"], + [{"weight": ((8,), "float32")}], + ] + sync = VllmRemoteSparseWeightSynchronizer( + policy, + generation, + transport="s3", + baseline_init_refs=[baseline_ref], + ) + + sync.init_communicator() + + policy.worker_group.run_all_workers_multiple_data.assert_not_called() + mock_ray.get.assert_any_call([baseline_ref]) + prepare.assert_called_once_with( + ["http://receiver"], + {"weight": ((8,), "float32")}, + api_key_env_var=None, + timeout_s=600.0, + ) + assert sync._baseline_init_refs == [] + + def test_merge_refit_info_rejects_conflicting_metadata(self): + with pytest.raises(ValueError, match="Conflicting sparse refit metadata"): + VllmRemoteSparseWeightSynchronizer._merge_refit_info( + [ + {"weight": ((8,), "float32")}, + {"weight": ((16,), "float32")}, + ] + ) + @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") def test_init_communicator_requires_receiver_endpoints(self, mock_ray): policy = MagicMock() @@ -141,7 +229,10 @@ def test_initializes_streams_commits_and_updates_baseline( assert [ entry.args[0] for entry in generation.worker_group.run_all_workers_single_data.call_args_list - ] == ["report_refit_server_base_url", "start_zmq_sparse_refit_relay"] + ] == [ + "report_refit_server_base_url", + "start_zmq_sparse_refit_relay", + ] generation.worker_group.run_all_workers_single_data.assert_any_call( "start_zmq_sparse_refit_relay", run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], diff --git a/tools/refit_bandwidth_calculator.py b/tools/refit_bandwidth_calculator.py index b4bcb0bc5e8..78f3b9772cd 100644 --- a/tools/refit_bandwidth_calculator.py +++ b/tools/refit_bandwidth_calculator.py @@ -32,9 +32,9 @@ tuple[Transport, Compression], tuple[tuple[float, float], tuple[float, float]] ] = { ("s3", "raw"): ((2.370, 116.138), (2.646, 183.431)), - ("s3", "zstd"): ((0.000, 84.746), (0.000, 134.041)), + ("s3", "zstd"): ((0.677, 84.672), (3.741, 145.597)), ("zmq", "raw"): ((0.000, 264.753), (0.000, 428.008)), - ("zmq", "zstd"): ((6.517, 73.621), (8.165, 157.339)), + ("zmq", "zstd"): ((5.813, 73.257), (0.000, 169.991)), } _NCCL_ANCHORS = ( (63.2, 0.84, 1.60), @@ -42,7 +42,7 @@ (470.2, 2.31, 2.73), (1342.0, 3.27, 3.46), ) -_WIRE_MULTIPLIER: dict[Compression, float] = {"raw": 2.0, "zstd": 0.74} +_WIRE_MULTIPLIER: dict[Compression, float] = {"raw": 2.0, "zstd": 0.747} @dataclass(frozen=True) From 1d37648501b37a1bbe6c2a4f0febdd340512119a Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Mon, 13 Jul 2026 12:49:38 -0700 Subject: [PATCH 09/18] Fix test cases Signed-off-by: Hollow Man --- tests/unit/models/megatron/test_community_import.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/tests/unit/models/megatron/test_community_import.py b/tests/unit/models/megatron/test_community_import.py index 21d18b03202..841e9dfe65d 100644 --- a/tests/unit/models/megatron/test_community_import.py +++ b/tests/unit/models/megatron/test_community_import.py @@ -30,7 +30,7 @@ def _ensure_package(monkeypatch, name: str) -> ModuleType: if "." in name: parent_name, child_name = name.rsplit(".", 1) parent_module = _ensure_package(monkeypatch, parent_name) - setattr(parent_module, child_name, module) + monkeypatch.setattr(parent_module, child_name, module, raising=False) return module @@ -87,14 +87,16 @@ def _install_runtime_stubs_for_hf_import(monkeypatch): parallel_state = ModuleType("megatron.core.parallel_state") parallel_state.model_parallel_is_initialized = lambda: False monkeypatch.setitem(sys.modules, "megatron.core.parallel_state", parallel_state) - core_module.parallel_state = parallel_state + monkeypatch.setattr(core_module, "parallel_state", parallel_state, raising=False) rerun_state_machine = ModuleType("megatron.core.rerun_state_machine") rerun_state_machine.destroy_rerun_state_machine = lambda: None monkeypatch.setitem( sys.modules, "megatron.core.rerun_state_machine", rerun_state_machine ) - core_module.rerun_state_machine = rerun_state_machine + monkeypatch.setattr( + core_module, "rerun_state_machine", rerun_state_machine, raising=False + ) tensor_parallel = ModuleType("megatron.core.tensor_parallel") tensor_parallel.model_parallel_cuda_manual_seed = lambda seed: None @@ -106,7 +108,7 @@ def _install_runtime_stubs_for_hf_import(monkeypatch): monkeypatch.setitem( sys.modules, "megatron.core.tensor_parallel.random", tensor_parallel_random ) - core_module.tensor_parallel = tensor_parallel + monkeypatch.setattr(core_module, "tensor_parallel", tensor_parallel, raising=False) def test_prefer_nvrx_is_noop_when_strategy_import_fails(monkeypatch): From 378c3293ddb5c89013ddbd592691d72caedaf097 Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Mon, 13 Jul 2026 21:28:51 -0700 Subject: [PATCH 10/18] Remove vllm internal overwrite Signed-off-by: Hollow Man --- docs/design-docs/sparse-delta-refit.md | 198 +++-- nemo_rl/algorithms/grpo.py | 12 +- .../models/generation/vllm/vllm_backend.py | 20 +- .../generation/vllm/vllm_sparse_delta.py | 786 +++++++++--------- .../generation/vllm/vllm_sparse_refit.py | 22 +- .../policy/workers/megatron_policy_worker.py | 8 +- .../workers/megatron_remote_sparse_refit.py | 243 +++--- nemo_rl/utils/weight_transfer_sparse_codec.py | 143 +--- .../generation/test_vllm_sparse_delta.py | 446 ++++------ .../generation/test_vllm_sparse_refit.py | 5 +- .../test_megatron_remote_sparse_refit.py | 104 +-- .../test_weight_transfer_remote_sparse.py | 59 +- 12 files changed, 854 insertions(+), 1192 deletions(-) diff --git a/docs/design-docs/sparse-delta-refit.md b/docs/design-docs/sparse-delta-refit.md index 0638cc2c413..21a0e3ac038 100644 --- a/docs/design-docs/sparse-delta-refit.md +++ b/docs/design-docs/sparse-delta-refit.md @@ -5,12 +5,12 @@ full checkpoint after every optimizer step. Megatron workers compare every uniquely owned MCore tensor against a policy-local CPU baseline. Exact affine mappings emit sparse Hugging Face (HF) coordinates directly; only changed tasks whose conversion is not affine traverse Megatron Bridge. S3 or ZeroMQ carries -the resulting payloads, and vLLM maps the HF coordinates into its local TP and -EP layout before applying them in place. +the resulting payloads, and each native vLLM weight loader applies its canonical +HF update to the local TP or EP destination. The feature is opt-in. Its synchronizer, codec, transports, receiver queue, and -placement engine are separate from existing NCCL, CUDA IPC, and packed refit -paths. +native-loader apply engine are separate from existing NCCL, CUDA IPC, and +packed refit paths. ## Supported scope @@ -51,7 +51,7 @@ flowchart LR subgraph G["vLLM generation cluster"] H["HTTP receiver"] Q["Eager node staging and bounded FIFO apply queue"] - A["Sparse placement and apply"] + A["Dense scratch and native-loader apply"] H --> Q --> A end @@ -59,8 +59,8 @@ flowchart LR E -->|ZeroMQ| Z --> H ``` -*Figure 1. S3 and ZeroMQ share the exporter, codec, receiver, placement engine, -and commit protocol.* +*Figure 1. S3 and ZeroMQ share the exporter, codec, receiver, native-loader +apply engine, and commit protocol.* | Responsibility | Implementation | |---|---| @@ -70,7 +70,7 @@ and commit protocol.* | Run the shared pipeline and S3 transport | [`weight_transfer_remote_sparse.py`](../../nemo_rl/utils/weight_transfer_remote_sparse.py) | | Run the ZeroMQ transport and relay | [`weight_transfer_zmq.py`](../../nemo_rl/utils/weight_transfer_zmq.py) | | Queue receiver work and expose endpoints | [`vllm_sparse_refit.py`](../../nemo_rl/models/generation/vllm/vllm_sparse_refit.py) | -| Map HF coordinates into vLLM tensors | [`vllm_sparse_delta.py`](../../nemo_rl/models/generation/vllm/vllm_sparse_delta.py) | +| Apply canonical updates through native loaders | [`vllm_sparse_delta.py`](../../nemo_rl/models/generation/vllm/vllm_sparse_delta.py) | ## Refit protocol @@ -106,9 +106,9 @@ residual Bridge export run concurrently. Baseline initialization also returns each canonical tensor's name, shape, and dtype. The synchronizer merges that metadata and asks every vLLM worker to -compile its native weight-loader placement plan before the first transfer. -This does not export HF values or mutate vLLM weights; it moves loader tracing -and plan validation out of the first timed refit. +reserve one reusable GPU byte buffer large enough for the largest canonical +tensor. This does not export HF values or mutate vLLM weights, and it removes +scratch allocation from the first timed refit. On a fresh run, vLLM already holds the shared checkpoint. Baseline construction starts early and can overlap initial generation, so the redundant initial full @@ -204,19 +204,26 @@ window at 256 payloads. Locations use `int32` unless a single canonical tensor exceeds the signed 32-bit index range; values remain grouped by dtype. The collective RPC passes only file paths, and workers use `torch.load(..., mmap=True)`. The mmap is not -a second baseline: it lets eight colocated ranks share the staged file's page -cache instead of materializing eight independent CPU copies, while each rank -copies only its selected entries to its GPU. - -Each worker selects the canonical entries consumed by its TP/EP placement plan -on CPU, converts only those selected locations to CUDA `int64`, and then copies -the selected locations and values to its GPU. Thus transport bytes are not -duplicated within a node and irrelevant canonical entries do not cross H2D. -The staged format flattens locations into one `int32` and one `int64` tensor; -it does not serialize one tensor object per model parameter. If ranks do not -share a node, the receiver still decodes once and sends that flat -representation through one collective RPC. Every worker derives and caches its -source plan from the native loader before CPU partitioning. +a second baseline: it lets colocated ranks share the staged file's page cache +instead of materializing independent CPU copies. The staged format flattens +locations into one `int32` and one `int64` tensor; it does not serialize one +tensor object per model parameter. If ranks do not share a node, the receiver +still decodes once and sends that flat representation through one collective +RPC. + +There is deliberately no TP/EP source plan. Every worker sees the canonical +sparse entries, scatters them into its reusable dense source buffer, and lets +the native loader select the local destination. This removes persistent vLLM +placement knowledge from NeMo RL, at the cost of canonical-tensor GPU +initialization and duplicated sparse H2D across ranks. Measure that cost on the +target TP/EP topology; H2D no longer scales only with the worker-local sparse +subset. + +During the existing untimed metadata prewarm, one no-op native-loader pass +records names that issue no model-storage copy on that fixed rank. Later refits +skip scratch construction for those explicit pipeline/expert/MTP skips. The +cache contains names only; it stores no placement offsets, tensor routes, or +model-family rules. The final `/nemo-rl/refit/flush` drains the queue, synchronizes CUDA, and checks optional delta samples. Only then does the source commit pending baseline @@ -228,7 +235,7 @@ updates in background CPU threads. > version before retrying. This is mandatory for `xor`, because replaying an > already-applied XOR reverts those bits. Replaying `overwrite` is safe. -## Payload and placement +## Payload and native apply Each serialized payload is: @@ -239,31 +246,34 @@ Each serialized payload is: Contiguous locations use a range encoding. Other sorted locations are delta-encoded into the smallest lossless unsigned width among 16, 32, and 64 bits. Metadata carries the HF name and shape, value offsets, location encoding, -the `xor` or `overwrite` operation, and optional verification samples. - -HF coordinates are the canonical wire format because Megatron Bridge already -defines the training-to-HF mapping while vLLM owns a different packed and -sharded layout. On first use, the receiver runs vLLM's native `load_weights()` -against metadata-only tensors while a PyTorch dispatch mode records the source -and destination views of each `copy_`. It caches those mappings and applies -later sparse values directly with bitwise XOR or `index_copy_`, without -materializing a dense HF tensor or duplicating QKV, MoE, Mamba, or TP placement -rules. - -The same recorded copies define a source plan for node-local partitioning. A -source plan contains the canonical offset, shape, and strides consumed by one -worker. Linear routes use sorted-range lookup; strided routes use affine source -coordinates. Verification samples are partitioned by the same plan. An unknown -or transformed loader fails while compiling the plan, before the receiver -stages or applies that tensor. - -The tracer accepts affine tensor views. Absolute overwrite also supports the -Mamba `A_log` transform and source-to-target dtype casts. Bridge residuals use -overwrite even when `encoding: xor` is configured because their target bits may -not match the HF source representation. XOR also rejects overlapping target -mappings. An element-expanding copy, unknown transform, or unplaced non-expert -tensor fails before any payload update. There is no dense fallback for an -unknown layout. +the `xor` or `overwrite` operation, and an optional verification sample budget. + +HF coordinates are the canonical wire format because Megatron Bridge defines +the training-to-HF mapping while vLLM owns the packed and sharded destination. +For each item, the receiver resets the resident largest-tensor scratch buffer, +scatters the sparse values, and calls the model's native `load_weights()`. +A storage-scoped PyTorch dispatch mode changes only copies into model parameter +or buffer storage; it never encodes QKV, MoE, Mamba, TP, or EP geometry. + +For XOR, unchanged scratch bits are zero and target copies become bitwise XOR. +The source must remain a view of the scratch storage, dtypes must match, and +overlapping destination copies fail closed. For overwrite, unchanged entries +are NaN sentinels. The dispatch mode propagates the first sparse mask through +subsequent native copies and writes only selected destination entries. It keeps +only those entries for per-item rollback rather than cloning the full target. +This supports native pointwise transforms and dtype casts without model-specific +formulas. One-byte FP8 overwrite uses an exact bit sentinel and therefore +requires a non-transforming native loader; end-to-end quantized rollout refit +remains out of scope. + +Native-loader return values distinguish an explicit skip from an unsupported +apply. An empty loaded set is accepted, matching vLLM's existing handling of +pipeline/expert ownership and checkpoint-only parameters such as inactive MTP +weights. Loader exceptions propagate. A loader that reports a weight loaded but +does not issue a supported target copy fails closed. There is no layout fallback +and no cached loader trace or route model. NeMo RL does not patch vLLM or encode +any vLLM layout: its only integration point is the model's public +`load_weights()` behavior. ## Configuration @@ -319,7 +329,7 @@ receiver apply. | Signal | Meaning | |---|---| | `REFIT_BASELINE_INIT` | Baseline export and snapshot time | -| `REFIT_RECEIVER_PREWARM` | Native vLLM placement plans compiled during initialization | +| `REFIT_RECEIVER_PREWARM` | GPU scratch reservation and rank-local native skip discovery | | `REFIT_{S3,ZMQ}_TIMING` | Producer wall time, stage service time, payloads, bytes, and changed density | | `REFIT_{S3,ZMQ}_DELTA_CHANGE` | Global changed and total element counts | | `REFIT_RECEIVER_TIMING` | Receiver staging span/wait, batches, apply time, and verification counts | @@ -329,7 +339,8 @@ receiver apply. `total_s` is producer wall time. Stage fields such as `encode_s`, `s3_put_s`, and `zmq_send_s` are sums across concurrent tasks and can exceed `total_s`; do not add them as serial phases. Receiver responses additionally expose node -decode/staging, worker deserialization, CPU partition, and sparse apply time. +decode/staging, worker deserialization, scratch preparation, and native-loader +apply time. These are also concurrent sums; compare them with receiver wall time rather than adding them. The `partition` field is `none` for uniquely owned policy-local shards, `names` for stable name-sharded residual exports, and @@ -353,18 +364,17 @@ sparse-apply NVTX ranges. Producer and receiver thread names begin with Keep transport changes behind the shared `stream_sparse_delta_payloads()` pipeline. A transport should provide payload delivery and timing only; it must -not duplicate the baseline tracker, codec, receiver queue, or placement logic. +not duplicate the baseline tracker, codec, receiver queue, or apply logic. Retries must preserve payload identity and bytes, fan out to every required replica, and require a successful global flush before baseline commit. Never retry XOR after an uncertain or partial receiver apply. -Do not add model-specific placement math. New layouts should work through their -native vLLM weight loader; extend the tracer only for a general loader operation -and fail closed for transformed or broadcasting copies. Tests must invoke the -real vLLM loader at nonzero TP ranks and cover replicated KV heads, packed -columns, local and remote experts, segmented views, contiguous ranges, and -explicit locations. Incorrect in-range XOR or overwrite locations silently -corrupt weights, so assert exact mapped indices and values. +Do not add model-specific placement math or a persistent placement cache. New +layouts must work through their native vLLM weight loader and the generic +storage-scoped operation context. Tests should cover packed QKV/MLP columns, +local and remote experts, segmented Mamba views, native transforms, dtype casts, +FP8 bit overwrite, contiguous ranges, and explicit locations. Unknown names, +transformed XOR, and overlapping XOR copies must fail closed. Codec changes must update encoder and decoder together, preserve 64-bit-safe locations, and commit exact source bits only after global success. Receiver @@ -412,41 +422,27 @@ the receiver before retrying. ## Refit bandwidth calculator [`refit_bandwidth_calculator.py`](../../tools/refit_bandwidth_calculator.py) is a -benchmark-specific estimator for the current S3 and ZeroMQ implementation. It -is not a general fabric or topology model. - -The zstd side embeds the July 12, 2026 end-to-end latency fits from 8 GB300 -sender GPUs in `us-east-2` to 64 H100 receiver GPUs in `us-east-1` shown above. -Measurements cover 63.2-1121.0 GB of indexed BF16 weights, S3 and ZeroMQ, and -3% and 5% changed density. The raw rows retain historical uncompressed -benchmark coefficients and are not a current production arm. Any positive -`--changed-pct` is accepted; values outside 3-5% use an extrapolated power -curve through the two measured-density fits. Estimated wire bytes use the -measured raw or zstd payload ratio. The coefficients in -`_SPARSE_LATENCY_FITS` implement -`fixed_seconds + seconds_per_1000_GB * model_size_GB / 1000`; they are latency -regressions, not bandwidth measurements. - -The NCCL side uses these measured generation-EP H100 refit envelopes on 400 -Gbps/rank InfiniBand: - -| Indexed BF16 | NCCL refit envelope | -|---:|---:| -| 63.2 GB | 0.84-1.60 s | -| 247.2 GB | 1.46-1.74 s | -| 470.2 GB | 2.31-2.73 s | -| 1342.0 GB | 3.27-3.46 s | - -The calculator interpolates these anchors in log model-size space, then -projects the full measured NCCL latency onto the candidate Ethernet rate: +calibrated comparison of the checked-in S3 and ZeroMQ measurements against a +measured H100 NCCL envelope. It is not a general fabric or topology simulator. + +The sparse side evaluates the latency fits in `_SPARSE_LATENCY_FITS` for the +requested model size, transport, compression, and any positive +`--changed-pct`. The 3% and 5% fits are measured calibration points; other +densities are explicit extrapolations. The coefficients model end-to-end +latency, not transport bandwidth, so `--candidate-ethernet-gbps` does not +rescale S3 or ZeroMQ. + +The NCCL side interpolates `_NCCL_ANCHORS` in log model-size space. Those +anchors were measured at 400 Gbps per rank and are projected onto the requested +Ethernet rate as: ```text T_ethernet = T_H100_IB * 400 / candidate_ethernet_gbps ``` -`--candidate-ethernet-gbps` is raw bandwidth per rank. It changes only the NCCL -projection; it does not rescale the measured S3 or ZeroMQ fit. Do not pass -aggregate node or cluster bandwidth. +`--candidate-ethernet-gbps` is raw bandwidth per rank, not aggregate node or +cluster bandwidth. This makes NCCL and the candidate Ethernet refer to the same +per-rank link while leaving the independently measured sparse path unchanged. ```bash uv run python tools/refit_bandwidth_calculator.py \ @@ -456,17 +452,19 @@ uv run python tools/refit_bandwidth_calculator.py \ --candidate-ethernet-gbps 25 ``` -The output reports the original H100 IB envelope, projected NCCL latency, -estimated sparse latency and wire bytes, and an Ethernet crossover range. Below -the lower crossover, sparse refit beats the complete NCCL envelope; above the -upper crossover, NCCL wins; between them, the measured NCCL range does not give -one winner. `--json` emits the same data for scripts. +The output reports the reference and projected NCCL envelopes, sparse latency, +estimated wire bytes, and the per-rank Ethernet crossover. Below the lower +crossover sparse refit beats the complete NCCL envelope; above the upper +crossover NCCL wins; between them the measured range has no single winner. +`--json` emits the same fields for scripts. The production transport currently applies zstd level 1 to every payload. -`--compression raw` selects a historical uncompressed benchmark fit for -analysis; it is not a runtime switch for the current transport. Treat model -sizes outside the measured range, changed densities outside 3-5%, and different -parallel mappings as experiment targets rather than performance claims. +`--compression raw` selects the checked-in uncompressed calibration for +analysis; it is not a runtime switch. Production payloads use zstd level 1. +Treat values outside the calibration range, or a different topology and +parallel mapping, as experiment inputs rather than performance claims. Update +the constants only from a balanced profile matrix and keep the source artifact +under `profiles/`. ## Failure guide @@ -474,7 +472,7 @@ parallel mappings as experiment targets rather than performance claims. |---|---| | Baseline is missing a tensor | Check baseline completion, checkpoint equality, and Bridge name mappings. | | No refit endpoint is found | Check worker startup, fixed ports, routing, and network policy. | -| No direct target plan exists | Add and unit-test the layout; do not silently fall back to dense loading. | +| Every worker reports a tensor unloaded | Verify the canonical HF name and native loader; do not add model-specific placement math. | | A payload ID is reused with different bytes | Start a new transfer or resend the original payload unchanged. | | Changed percentage rises unexpectedly | Correlate `DELTA_CHANGE` with `GLOBAL_COMMIT` and baseline commit completion. | | Apply queue stalls | Inspect receiver timing and reduce source or relay concurrency. | diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index d716bc77102..7292c40fa5d 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -1106,14 +1106,12 @@ def initialize_generation_with_policy( validate_vllm_remote_sparse_refit, ) - remote_transport = cast( - str, - validate_vllm_remote_sparse_refit( - generation_config, - colocated=colocated_inference, - megatron_enabled=policy_config["megatron_cfg"]["enabled"], - ), + remote_transport = validate_vllm_remote_sparse_refit( + generation_config, + colocated=colocated_inference, + megatron_enabled=policy_config["megatron_cfg"]["enabled"], ) + assert remote_transport is not None def init_policy_for_generation(): policy, policy_time = init_policy() diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index 3e352d5fa2e..9a3e76d72ef 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -176,8 +176,10 @@ def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: def prepare_sparse_delta_refit_info( self, state_dict_info: dict[str, tuple[tuple[int, ...], torch.dtype]] ) -> None: - """Compile sparse placement plans before the first timed refit.""" - self._get_sparse_delta_applier().prewarm(state_dict_info) + """Reserve the reusable sparse-refit scratch buffer.""" + applier = self._get_sparse_delta_applier() + applier.prewarm(state_dict_info) + applier.discover_native_skips(state_dict_info) def _maybe_process_fp8_kv_cache(self) -> None: """Process weights after loading for FP8 KV cache (static scales).""" @@ -370,7 +372,7 @@ def _load_weights(self, weights): def _get_sparse_delta_applier(self) -> Any: if self._sparse_delta_applier is None: - # Avoid importing sparse placement code for existing refit transports. + # Avoid importing sparse-refit code for existing refit transports. from nemo_rl.models.generation.vllm.vllm_sparse_delta import ( VllmSparseDeltaApplier, ) @@ -514,19 +516,15 @@ def update_weights_from_decoded_sparse_payload( self, *serialized_payloads: bytes, ) -> dict[str, Any]: - return ( - self._get_sparse_delta_applier().update_weights_from_decoded_sparse_payload( - *serialized_payloads - ) - ) + applier = self._get_sparse_delta_applier() + return applier.update_weights_from_decoded_sparse_payload(*serialized_payloads) def update_weights_from_decoded_sparse_payload_files( self, *payload_paths: str, ) -> dict[str, Any]: - return self._get_sparse_delta_applier().update_weights_from_decoded_sparse_payload_files( - *payload_paths - ) + applier = self._get_sparse_delta_applier() + return applier.update_weights_from_decoded_sparse_payload_files(*payload_paths) def synchronize_device(self) -> None: self._get_sparse_delta_applier().synchronize_device() diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py index f236e32d4b8..aa2c4277a81 100644 --- a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py +++ b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py @@ -12,14 +12,11 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Direct sparse-delta placement and application for vLLM workers.""" +"""Apply canonical sparse updates through vLLM's native weight loaders.""" import io -import re import time -from collections import defaultdict -from collections.abc import Mapping -from dataclasses import dataclass +from collections.abc import Iterable, Iterator, Mapping from math import prod from typing import Any, cast @@ -29,61 +26,84 @@ from nemo_rl.utils import weight_transfer_sparse_codec as sparse_codec from nemo_rl.utils.nsys import wrap_with_nvtx_name -_EXPERT_WEIGHT_RE = re.compile(r"\.experts\.\d+\.(?:gate|up|down)_proj\.weight$") - def _storage_key(tensor: torch.Tensor) -> int: return tensor.untyped_storage()._cdata -@dataclass(frozen=True) -class _SparseDeltaCopyPlan: - target: torch.Tensor - source_offset: int - source_strides: tuple[int, ...] - shape: tuple[int, ...] - target_offset: int - target_strides: tuple[int, ...] - linear: bool - - -@dataclass(frozen=True) -class _SparseDeltaTargetPlan: - copies: tuple[_SparseDeltaCopyPlan, ...] = () - log_delta_transform: bool = False - identity: bool = False +def _integer_view(tensor: torch.Tensor) -> torch.Tensor: + return tensor.view( + sparse_codec.integer_dtype_for_element_size(tensor.element_size()) + ) -class _SparseLoadTracer(TorchDispatchMode): - """Capture the views copied by vLLM's native weight loaders.""" +class _SparseWeightLoadMode(TorchDispatchMode): + """Turn native loader copies into sparse XOR or transactional overwrite.""" def __init__( self, - targets: list[torch.Tensor], - sources: dict[str, torch.Tensor], + targets: set[int], + verification: list[tuple[torch.Tensor, torch.Tensor, torch.Tensor]], ) -> None: super().__init__() - self.copies: dict[str, list[_SparseDeltaCopyPlan]] = defaultdict(list) - self.postprocessed: set[str] = set() - self._sources = {_storage_key(tensor): name for name, tensor in sources.items()} - self._targets: dict[int, list[torch.Tensor]] = defaultdict(list) - for target in targets: - if not target.is_contiguous(): - raise RuntimeError("Sparse delta targets must be contiguous.") - self._targets[_storage_key(target)].append(target) - self._last_source: dict[int, str] = {} - - def _target_for(self, view: torch.Tensor) -> torch.Tensor | None: - candidates = self._targets.get(_storage_key(view), ()) - view_start = view.storage_offset() - view_end = view_start + sum( - (size - 1) * stride for size, stride in zip(view.shape, view.stride()) + self._targets = targets + self._verification = verification + self._source_storage = 0 + self._operation: sparse_codec.SparseOperation = "overwrite" + self._sample_limit = 0 + self._exact_sentinel: int | None = None + self._backups: list[tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = [] + self._active_masks: dict[ + tuple[int, int, tuple[int, ...], tuple[int, ...]], torch.Tensor + ] = {} + self._verification_masks: dict[ + tuple[int, int, tuple[int, ...], tuple[int, ...]], + tuple[torch.Tensor, torch.Tensor], + ] = {} + self._xor_spans: dict[int, list[tuple[int, int]]] = {} + self.copies = 0 + + def start( + self, + source: torch.Tensor, + operation: sparse_codec.SparseOperation, + sample_limit: int, + exact_sentinel: int | None, + ) -> None: + if self._backups: + raise RuntimeError("Previous sparse native-loader item was not finished.") + self._source_storage = _storage_key(source) + self._operation = operation + self._sample_limit = sample_limit + self._exact_sentinel = exact_sentinel + self._active_masks.clear() + self._verification_masks.clear() + self._xor_spans.clear() + self.copies = 0 + + def _overwrite_changed( + self, + destination: torch.Tensor, + changed: torch.Tensor, + values: torch.Tensor, + ) -> None: + backup = destination.masked_select(changed) + self._backups.append((destination, changed, backup)) + destination.masked_scatter_(changed, values.to(destination.dtype)) + + def _remember_changed( + self, destination: torch.Tensor, changed: torch.Tensor + ) -> None: + view_key = ( + _storage_key(destination), + int(destination.storage_offset()), + tuple(int(size) for size in destination.shape), + tuple(int(stride) for stride in destination.stride()), ) - for target in candidates: - start = target.storage_offset() - if start <= view_start and view_end < start + target.numel(): - return target - return None + previous = self._verification_masks.get(view_key) + if previous is not None: + changed = previous[1] | changed + self._verification_masks[view_key] = (destination, changed) def __torch_dispatch__( self, @@ -96,331 +116,333 @@ def __torch_dispatch__( return func(*args, **(kwargs or {})) destination, source = cast(tuple[torch.Tensor, torch.Tensor], args[:2]) - target = self._target_for(destination) - if target is None: - raise RuntimeError("vLLM loader copied outside a model parameter.") - - source_name = self._sources.get(_storage_key(source)) - target_key = id(target) - if source_name is None: - source_name = self._last_source.get(target_key) - if source_name is None: - raise RuntimeError("vLLM loader materialized an unsupported transform.") - self.postprocessed.add(source_name) + if _storage_key(destination) not in self._targets: + return func(*args, **(kwargs or {})) + + self.copies += 1 + if self._operation == "overwrite": + if not source.dtype.is_floating_point: + raise RuntimeError("Sparse overwrite requires a floating-point loader.") + source = source.expand_as(destination) + view_key = ( + _storage_key(destination), + int(destination.storage_offset()), + tuple(int(size) for size in destination.shape), + tuple(int(stride) for stride in destination.stride()), + ) + if self._exact_sentinel is not None: + if ( + _storage_key(source) != self._source_storage + or source.dtype != destination.dtype + ): + raise RuntimeError( + "Exact FP8 overwrite cannot pass through a transforming loader." + ) + source_bits = _integer_view(source) + changed = source_bits.ne(self._exact_sentinel) + destination_bits = _integer_view(destination) + self._overwrite_changed( + destination_bits, + changed, + source_bits.masked_select(changed), + ) + self._remember_changed(destination, changed) + return destination + changed = self._active_masks.get(view_key) + if changed is None or _storage_key(source) == self._source_storage: + changed = ~torch.isnan(source) + self._active_masks[view_key] = changed + self._overwrite_changed( + destination, + changed, + source.masked_select(changed), + ) + self._remember_changed(destination, changed) return destination - if destination.shape != source.shape: - raise RuntimeError("vLLM loader used an expanding copy.") - shape = tuple(source.shape) - source_strides = tuple(source.stride()) - target_strides = tuple(destination.stride()) - if any(size > 1 and stride <= 0 for size, stride in zip(shape, source_strides)): - raise RuntimeError("vLLM loader used an unsupported source view.") - contiguous = torch.empty(shape, device="meta").stride() - self.copies[source_name].append( - _SparseDeltaCopyPlan( - target, - int(source.storage_offset()), - source_strides, - shape, - int(destination.storage_offset() - target.storage_offset()), - target_strides, - source_strides == target_strides == contiguous, + if _storage_key(source) != self._source_storage: + raise RuntimeError( + "XOR cannot pass through a native loader that transforms its input." ) + if source.dtype != destination.dtype: + raise RuntimeError("XOR source and target dtypes must match.") + source = source.expand_as(destination) + origin = int(destination.storage_offset()) + extents = [ + (int(size) - 1) * int(stride) + for size, stride in zip( + destination.shape, destination.stride(), strict=True + ) + ] + span = ( + origin + sum(min(0, extent) for extent in extents), + origin + sum(max(0, extent) for extent in extents), + ) + spans = self._xor_spans.setdefault(_storage_key(destination), []) + if any(span[0] <= other[1] and other[0] <= span[1] for other in spans): + raise RuntimeError("XOR native loader produced overlapping target copies.") + spans.append(span) + destination_bits = _integer_view(destination) + source_bits = _integer_view(source) + changed = source_bits.ne(0) + values = destination_bits.masked_select(changed).bitwise_xor( + source_bits.masked_select(changed) ) - self._last_source[target_key] = source_name + self._overwrite_changed(destination_bits, changed, values) + self._remember_changed(destination, changed) return destination + def _record(self, target: torch.Tensor, changed: torch.Tensor) -> None: + if self._sample_limit <= 0 or not target.is_contiguous(): + return + locations = changed.reshape(-1).nonzero().reshape(-1)[: self._sample_limit] + if not locations.numel(): + return + target_bits = _integer_view(target).reshape(-1) + self._verification.append( + (target, locations, target_bits.index_select(0, locations).clone()) + ) + self._sample_limit -= locations.numel() + + @torch.no_grad() + def finish(self) -> None: + """Commit the sparse copies after the loader finishes all transforms.""" + for destination, changed in self._verification_masks.values(): + self._record(destination, changed) + self._backups.clear() + self._verification_masks.clear() + + @torch.no_grad() + def rollback(self) -> None: + for destination, changed, backup in reversed(self._backups): + destination.masked_scatter_(changed, backup) + self._backups.clear() + self._verification_masks.clear() + class VllmSparseDeltaApplier: - """Apply sparse HF deltas through plans derived from native vLLM loaders.""" + """Own one dense GPU scratch buffer and delegate all placement to vLLM.""" def __init__(self, model_runner: Any, device: torch.device) -> None: self.model_runner = model_runner self._cuda_device_index = device.index - self._plan_cache: dict[str, _SparseDeltaTargetPlan] = {} + model = model_runner.model + self._target_storages = { + _storage_key(tensor) for tensor in (*model.parameters(), *model.buffers()) + } + self._scratch = torch.empty(0, dtype=torch.uint8, device=device) + self._classified_names: set[str] = set() + self._skipped_names: set[str] = set() self._verification: list[tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = [] self._verification_candidates = 0 - def _compile_plans(self, metadata: list[dict[str, Any]]) -> None: - missing = { - str(item["name"]): item - for item in metadata - if item["name"] not in self._plan_cache - } - if not missing: - return - - model = self.model_runner.model - targets = list(model.parameters()) + list(model.buffers()) - sources = { - name: torch.empty( - tuple(item["shape"]), - dtype=sparse_codec.dtype_from_name(str(item["dtype"])), - device="meta", - ) - for name, item in missing.items() - } - tracer = _SparseLoadTracer(targets, sources) - with torch.no_grad(), tracer: - model.load_weights((name, sources[name]) for name in missing) - - for name, item in missing.items(): - source_shape = tuple(item["shape"]) - source_strides = tuple(sources[name].stride()) - copies = tuple(tracer.copies.get(name, ())) - transformed = name in tracer.postprocessed - log_transform = ( - transformed and ".mixer." in name and name.endswith((".A", ".A_log")) - ) - if transformed and not log_transform: - raise RuntimeError( - f"vLLM loader for {name!r} transforms weights and cannot apply deltas." - ) - if not copies and not ( - name.startswith(("mtp.", "draft.")) or _EXPERT_WEIGHT_RE.search(name) - ): - raise RuntimeError(f"vLLM loader did not place {name!r}.") - identity = ( - len(copies) == 1 - and not log_transform - and copies[0].source_offset == 0 - and copies[0].shape == source_shape - and copies[0].source_strides == source_strides - and copies[0].target_offset == 0 - and copies[0].target_strides == source_strides - and copies[0].target.numel() == prod(source_shape) - ) - self._plan_cache[name] = _SparseDeltaTargetPlan( - copies, log_transform, identity - ) - - def sparse_delta_source_plans( - self, metadata: list[dict[str, Any]] - ) -> dict[str, sparse_codec.SparseSourcePlan]: - """Describe which canonical source views this worker consumes.""" - self._compile_plans(metadata) - return { - name: sparse_codec.SparseSourcePlan( - routes=tuple( - sparse_codec.SparseSourceRoute( - copy.source_offset, - copy.source_strides, - copy.shape, - copy.linear, - ) - for copy in plan.copies - ), - identity=plan.identity, - ) - for name, plan in ( - (str(item["name"]), self._plan_cache[str(item["name"])]) - for item in metadata - ) - } - def prewarm( self, state_dict_info: Mapping[str, tuple[tuple[int, ...], torch.dtype]] ) -> None: - self._compile_plans( - [ - { - "name": name, - "shape": shape, - "dtype": str(dtype).removeprefix("torch."), - } - for name, (shape, dtype) in state_dict_info.items() - ] - ) - - @staticmethod - def _map_copy( - locations: torch.Tensor, - values: torch.Tensor, - copy: _SparseDeltaCopyPlan, - ) -> tuple[torch.Tensor, torch.Tensor]: - mapped, keep = sparse_codec.map_sparse_locations( - locations, - copy.source_offset, - copy.source_strides, - copy.shape, - copy.linear, - copy.target_offset, - copy.target_strides, + """Reserve a reusable buffer for the largest canonical source tensor.""" + required = max( + (prod(shape) * dtype.itemsize for shape, dtype in state_dict_info.values()), + default=0, ) - return mapped[keep], values[keep] + if required > self._scratch.numel(): + self._scratch = torch.empty( + required, dtype=torch.uint8, device=self._scratch.device + ) - def _record_verification( - self, - item: dict[str, Any], - plan: _SparseDeltaTargetPlan, - operation: sparse_codec.SparseOperation, - source_dtype: torch.dtype, + def discover_native_skips( + self, state_dict_info: Mapping[str, tuple[tuple[int, ...], torch.dtype]] ) -> None: - sample_locations = item.get("verification_locations", []) - self._verification_candidates += len(sample_locations) - if not sample_locations or not plan.copies: + """Cache weights that the native loader explicitly skips on this rank.""" + pending = [ + (name, shape, dtype) + for name, (shape, dtype) in state_dict_info.items() + if name not in self._classified_names and dtype.is_floating_point + ] + if not pending: return - target = plan.copies[0].target - locations = torch.tensor(sample_locations, device=target.device) - value_dtype = sparse_codec.integer_dtype_for_element_size(source_dtype.itemsize) - values = torch.tensor( - item["verification_values"], device=target.device, dtype=value_dtype - ) - for copy in plan.copies: - mapped, selected = self._map_copy(locations, values, copy) - if not mapped.numel(): - continue - if operation == "xor": - target_bits = self._integer_flat(copy.target) - expected = target_bits.index_select(0, mapped).bitwise_xor(selected) - else: - _, replacement = self._overwrite_target_values( - copy.target, - selected, - source_dtype, - log_transform=plan.log_delta_transform, + + mode = _SparseWeightLoadMode(self._target_storages, []) + observations: list[tuple[str, int]] = [] + + def weights() -> Iterator[tuple[str, torch.Tensor]]: + active_name = None + for name, shape, dtype in pending: + if active_name is not None: + mode.finish() + observations.append((active_name, mode.copies)) + source = self._source_tensor( + { + "name": name, + "shape": shape, + "dtype": str(dtype).removeprefix("torch."), + } ) - expected = replacement.contiguous().view( - sparse_codec.integer_dtype_for_element_size( - copy.target.element_size() - ) + source.fill_(float("nan")) + exact_sentinel = ( + int(_integer_view(source).reshape(-1)[0].item()) + if source.element_size() == 1 + else None ) - self._verification.append((copy.target, mapped, expected)) + mode.start(source, "overwrite", 0, exact_sentinel) + active_name = name + yield name, source + if active_name is not None: + mode.finish() + observations.append((active_name, mode.copies)) + + loader_weights = weights() + try: + with torch.no_grad(), mode: + loaded = self.model_runner.model.load_weights(loader_weights) + except Exception: + mode.rollback() + raise + finally: + loader_weights.close() + + if len(observations) != len(pending): + raise RuntimeError( + "Native loader did not consume all sparse weight metadata." + ) + copied = sum(copies > 0 for _, copies in observations) + if loaded is not None and len(loaded) > copied: + raise RuntimeError( + "Native loader reported a loaded sparse weight without a " + "supported target copy." + ) + self._classified_names.update(name for name, _ in observations) + if loaded is not None: + self._skipped_names.update( + name for name, copies in observations if copies == 0 + ) - @staticmethod - def _integer_flat(target: torch.Tensor) -> torch.Tensor: - dtype = sparse_codec.integer_dtype_for_element_size(target.element_size()) - return target.data.view(dtype).view(-1) + def _source_tensor(self, item: dict[str, Any]) -> torch.Tensor: + shape = tuple(int(dim) for dim in item["shape"]) + dtype = sparse_codec.dtype_from_name(str(item["dtype"])) + byte_count = prod(shape) * dtype.itemsize + if byte_count > self._scratch.numel(): + self.prewarm({str(item["name"]): (shape, dtype)}) + return self._scratch[:byte_count].view(dtype).view(shape) @staticmethod - def _xor_target_mappings_overlap(plan: _SparseDeltaTargetPlan) -> bool: - spans: dict[int, list[tuple[int, int]]] = defaultdict(list) - for copy in plan.copies: - origin = int(copy.target.storage_offset()) + copy.target_offset - extents = [ - (size - 1) * stride - for size, stride in zip(copy.shape, copy.target_strides, strict=True) - ] - start = origin + sum(min(0, extent) for extent in extents) - end = origin + sum(max(0, extent) for extent in extents) - target_spans = spans[_storage_key(copy.target)] - if any( - start <= other_end and other_start <= end - for other_start, other_end in target_spans - ): - return True - target_spans.append((start, end)) - return False - - @classmethod - def _overwrite_target_values( - cls, - target: torch.Tensor, - values: torch.Tensor, - source_dtype: torch.dtype, - *, - log_transform: bool, - ) -> tuple[torch.Tensor, torch.Tensor]: - if not log_transform and target.dtype == source_dtype: - return cls._integer_flat(target), values - source_values = values.contiguous().view(source_dtype) - replacement = ( - -source_values.float().exp().to(target.dtype) - if log_transform - else source_values.to(target.dtype) - ) - return target.data.view(-1), replacement - - def _apply_decoded_item( - self, + def _scatter_values( + source: torch.Tensor, item: dict[str, Any], - plan: _SparseDeltaTargetPlan, locations: torch.Tensor, values: torch.Tensor, ) -> None: - operation = sparse_codec.sparse_operation(item["operation"]) - source_dtype = sparse_codec.dtype_from_name(str(item["dtype"])) - if not plan.copies: - self._record_verification(item, plan, operation, source_dtype) - return - first_target = plan.copies[0].target - if operation == "xor" and plan.log_delta_transform: - raise RuntimeError(f"XOR cannot apply transformed weight {item['name']!r}.") - if operation == "xor" and any( - copy.target.dtype != source_dtype for copy in plan.copies - ): - raise RuntimeError( - f"XOR source and target dtypes differ for {item['name']!r}." - ) - if operation == "xor" and self._xor_target_mappings_overlap(plan): - raise RuntimeError(f"XOR target mappings overlap for {item['name']!r}.") + source_bits = _integer_view(source).reshape(-1) expected_dtype = sparse_codec.integer_dtype_for_element_size( - source_dtype.itemsize + source.element_size() ) if values.dtype != expected_dtype: raise RuntimeError( f"Sparse values have the wrong dtype for {item['name']!r}." ) - values = values.to(device=first_target.device, non_blocking=True) - self._record_verification(item, plan, operation, source_dtype) - if plan.identity and item["index_encoding"] == "range": - if operation == "xor": - target = self._integer_flat(first_target) - target.narrow(0, int(item["range_start"]), values.numel()).bitwise_xor_( - values - ) - else: - target, replacement = self._overwrite_target_values( - first_target, values, source_dtype, log_transform=False - ) - target.narrow(0, int(item["range_start"]), values.numel()).copy_( - replacement - ) + values = values.to(device=source.device, non_blocking=True) + if item["index_encoding"] == "range": + source_bits.narrow(0, int(item["range_start"]), values.numel()).copy_( + values + ) return + source_bits.index_copy_( + 0, + locations.to(device=source.device, dtype=torch.int64, non_blocking=True), + values, + ) - locations = locations.to( - device=first_target.device, dtype=torch.int64, non_blocking=True + def _prepare_loader_weight( + self, + item: dict[str, Any], + locations: torch.Tensor, + values: torch.Tensor, + mode: _SparseWeightLoadMode, + ) -> tuple[str, torch.Tensor]: + operation = sparse_codec.sparse_operation(item["operation"]) + source = self._source_tensor(item) + exact_sentinel = None + if operation == "xor": + source.zero_() + elif not source.dtype.is_floating_point: + raise RuntimeError("Sparse overwrite requires a floating-point source.") + else: + source.fill_(float("nan")) + if source.element_size() == 1: + source_bits = _integer_view(source) + exact_sentinel = int(source_bits.reshape(-1)[0].item()) + if bool(values.eq(exact_sentinel).any()): + exact_sentinel ^= 0x80 + if bool(values.eq(exact_sentinel).any()): + raise RuntimeError( + "FP8 sparse overwrite exhausted its sentinel values." + ) + source_bits.fill_(exact_sentinel) + self._scatter_values(source, item, locations, values) + + sample_limit = int(item.get("verification_samples", 0)) + self._verification_candidates += sample_limit + mode.start( + source, + operation, + sample_limit, + exact_sentinel, ) - if plan.identity: - if operation == "xor": - target = self._integer_flat(first_target) - current = target.index_select(0, locations) - target.index_copy_(0, locations, current.bitwise_xor(values)) - else: - target, replacement = self._overwrite_target_values( - first_target, values, source_dtype, log_transform=False + return str(item["name"]), source + + def _apply_decoded_items( + self, + items: Iterable[tuple[dict[str, Any], torch.Tensor, torch.Tensor]], + ) -> None: + mode = _SparseWeightLoadMode(self._target_storages, self._verification) + yielded_names: list[str] = [] + copy_counts: list[int] = [] + + def weights() -> Iterator[tuple[str, torch.Tensor]]: + active = False + for item, locations, values in items: + if str(item["name"]) in self._skipped_names: + continue + if active: + mode.finish() + copy_counts.append(mode.copies) + weight = self._prepare_loader_weight(item, locations, values, mode) + active = True + yielded_names.append(weight[0]) + yield weight + if active: + mode.finish() + copy_counts.append(mode.copies) + + loader_weights = weights() + try: + with torch.no_grad(), mode: + loaded = self.model_runner.model.load_weights(loader_weights) + if len(copy_counts) != len(yielded_names): + raise RuntimeError("Native loader did not consume all sparse weights.") + copied_items = sum(copies > 0 for copies in copy_counts) + if loaded is None and copied_items != len(yielded_names): + raise RuntimeError( + "Native loader did not report whether uncopied sparse weights " + "were skipped." ) - target.index_copy_(0, locations, replacement) - return - for copy in plan.copies: - mapped, selected = self._map_copy(locations, values, copy) - if not mapped.numel(): - continue - if operation == "xor": - target = self._integer_flat(copy.target) - current = target.index_select(0, mapped) - target.index_copy_(0, mapped, current.bitwise_xor(selected)) - else: - target, replacement = self._overwrite_target_values( - copy.target, - selected, - source_dtype, - log_transform=plan.log_delta_transform, + if loaded is not None and len(loaded) > copied_items: + raise RuntimeError( + "Native loader reported a loaded sparse weight without a " + "supported target copy." ) - target.index_copy_(0, mapped, replacement) + except Exception: + mode.rollback() + raise + finally: + loader_weights.close() - def _apply_decoded_sparse_weight_deltas( - self, decoded: list[sparse_codec.DecodedSparseItem] + def _apply_decoded_item( + self, + item: dict[str, Any], + locations: torch.Tensor, + values: torch.Tensor, ) -> None: - with torch.no_grad(): - for item, locations, values in decoded: - self._apply_decoded_item( - item, - self._plan_cache[str(item["name"])], - locations, - values, - ) + self._apply_decoded_items(((item, locations, values),)) @wrap_with_nvtx_name( "vllm_internal_worker_extension/update_weights_from_decoded_sparse_payload" @@ -454,24 +476,28 @@ def _load_decoded_sparse_payloads( deserialize_s += time.perf_counter() - item_started item_started = time.perf_counter() - metadata = [item for _, _, items in payloads for item in items] - source_plans = self.sparse_delta_source_plans(metadata) - plan_s = time.perf_counter() - item_started - partition_s = sparse_apply_s = 0.0 - for payload in payloads: - item_started = time.perf_counter() - selected = sparse_codec.partition_decoded_sparse_entries( - sparse_codec.iter_decoded_sparse_payload(payload), source_plans - ) - partition_s += time.perf_counter() - item_started - item_started = time.perf_counter() - self._apply_decoded_sparse_weight_deltas(selected) - sparse_apply_s += time.perf_counter() - item_started + self.prewarm( + { + str(item["name"]): ( + tuple(item["shape"]), + sparse_codec.dtype_from_name(str(item["dtype"])), + ) + for _, _, items in payloads + for item in items + } + ) + scratch_s = time.perf_counter() - item_started + item_started = time.perf_counter() + self._apply_decoded_items( + decoded_item + for payload in payloads + for decoded_item in sparse_codec.iter_decoded_sparse_payload(payload) + ) + sparse_apply_s = time.perf_counter() - item_started return { "ok": True, "receiver_deserialize_s": deserialize_s, - "receiver_plan_s": plan_s, - "receiver_partition_s": partition_s, + "receiver_scratch_s": scratch_s, "receiver_sparse_apply_s": sparse_apply_s, "receiver_total_s": time.perf_counter() - started, } @@ -486,60 +512,46 @@ def synchronize_device(self) -> None: torch.cuda.synchronize(self._cuda_device_index) def finish_sparse_delta_refit(self) -> dict[str, Any]: - """Synchronize and compare bounded producer samples with applied weights.""" + """Synchronize and compare bounded samples of target entries just changed.""" self.synchronize_device() verification, self._verification = self._verification, [] candidates, self._verification_candidates = self._verification_candidates, 0 + stats = ( + torch.zeros(4, device=verification[0][0].device) if verification else None + ) samples = 0 - stats = [0.0] * 4 - if verification: - with torch.no_grad(): - differences = [] - exact_mismatches = [] - mismatches = [] - for target, locations, expected_bits in verification: - integer_dtype = sparse_codec.integer_dtype_for_element_size( - target.element_size() - ) - actual_bits = ( - target.data.view(integer_dtype) - .view(-1) - .index_select(0, locations) - ) - bit_mismatches = actual_bits.ne(expected_bits) - actual = target.data.view(-1).index_select(0, locations).float() - expected = expected_bits.view(target.dtype).float() - difference = torch.where( + with torch.no_grad(): + for target, locations, expected_bits in verification: + actual_bits = ( + _integer_view(target).reshape(-1).index_select(0, locations) + ) + bit_mismatches = actual_bits.ne(expected_bits) + actual = actual_bits.view(target.dtype).float() + expected = expected_bits.view(target.dtype).float() + difference = torch.nan_to_num( + torch.where( bit_mismatches, (actual - expected).abs(), torch.zeros_like(actual), - ) - differences.append(torch.nan_to_num(difference, nan=float("inf"))) - exact_mismatches.append(bit_mismatches) - mismatches.append( - bit_mismatches - & ~torch.isclose(actual, expected, rtol=1e-6, atol=1e-8) - ) - samples += actual.numel() - difference = torch.cat(differences) - stats = ( - torch.stack( - ( - difference.sum(), - difference.max(), - torch.cat(exact_mismatches).sum().float(), - torch.cat(mismatches).sum().float(), - ) - ) - .cpu() - .tolist() + ), + nan=float("inf"), ) + assert stats is not None + stats[0] += difference.sum() + stats[1] = torch.maximum(stats[1], difference.max()) + stats[2] += bit_mismatches.sum() + stats[3] += ( + bit_mismatches + & ~torch.isclose(actual, expected, rtol=1e-6, atol=1e-8) + ).sum() + samples += actual.numel() + values = [0.0] * 4 if stats is None else stats.cpu().tolist() return { "ok": True, "verification_candidates": candidates, "verification_samples": samples, - "verification_exact_mismatches": int(stats[2]), - "verification_mismatches": int(stats[3]), - "verification_abs_sum": float(stats[0]), - "verification_max_abs": float(stats[1]), + "verification_exact_mismatches": int(values[2]), + "verification_mismatches": int(values[3]), + "verification_abs_sum": float(values[0]), + "verification_max_abs": float(values[1]), } diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py index 41ad65cc754..0d52c21fb05 100644 --- a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py +++ b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py @@ -65,7 +65,7 @@ def _decode_staged_payload( ) return ( sparse_codec.decode_sparse_tensor_payload_for_staging(payload), - sum(len(item.get("verification_locations", ())) for item in payload[2]), + sum(int(item.get("verification_samples", 0)) for item in payload[2]), ) @@ -205,16 +205,12 @@ def _enqueue_sparse_payload_apply( def _submit_pending_sparse_payloads(self) -> Future[dict[str, Any]]: payloads = tuple(self._refit_apply_pending_payloads) self._refit_apply_pending_payloads.clear() - if self._refit_workers_share_node: - future = self._refit_apply_executor.submit( - self.update_weights_from_staged_sparse_payloads, - cast(tuple[Future[_StagedSparsePayload], ...], payloads), - ) - else: - future = self._refit_apply_executor.submit( - self.update_weights_from_serialized_sparse_payloads, - cast(tuple[bytes, ...], payloads), - ) + apply = ( + self.update_weights_from_staged_sparse_payloads + if self._refit_workers_share_node + else self.update_weights_from_serialized_sparse_payloads + ) + future = self._refit_apply_executor.submit(cast(Any, apply), payloads) self._refit_apply_futures.append(future) future.add_done_callback(self._notify_refit_apply_waiters) return future @@ -228,12 +224,10 @@ def _collect_refit_apply_results( futures: list[Future[dict[str, Any]]], ) -> dict[str, Any]: results = [future.result() for future in futures] - timing: dict[str, float] = {} - merge_vllm_refit_metrics(timing, results, maximum=False) return { "ok": True, "payloads": sum(int(result.get("payloads", 0)) for result in results), - **timing, + **merge_vllm_refit_metrics({}, results, maximum=False), } @staticmethod diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 209e06371f7..48bef53d731 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -1915,14 +1915,12 @@ def _iter_params_with_optional_kv_scales( get_vllm_qkv_scale_names, ) + if conversion_tasks is None: + conversion_tasks = self.refit_conversion_tasks base_iter = self.megatron_bridge.export_hf_weights( [self.model], show_progress=False, - conversion_tasks=( - self.refit_conversion_tasks - if conversion_tasks is None - else conversion_tasks - ), + conversion_tasks=conversion_tasks, ) # Yield the original parameters first. diff --git a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py index 647b1100161..398836cfe0e 100644 --- a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py +++ b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py @@ -18,7 +18,7 @@ from collections.abc import Iterable, Mapping from concurrent.futures import ThreadPoolExecutor from functools import cache, partial -from typing import Any +from typing import Any, cast import torch @@ -58,17 +58,15 @@ def __init__(self, worker: Any, delta_config: Mapping[str, Any]) -> None: if residual_config["encoding"] == "xor": residual_config["encoding"] = "overwrite" self._tracker = DeltaCompressionTracker(residual_config) - self._local_tracker: DeltaCompressionTracker | None = None - self._change_tracker: DeltaCompressionTracker | None = None + self._policy_tracker: DeltaCompressionTracker | None = None self._local_tensors: list[tuple[str, torch.Tensor]] = [] self._misc_local_tensors: list[tuple[str, torch.Tensor]] = [] self._misc_conversion_tasks: list[Any] | None = None self._filter_misc_tasks = False - self._uses_policy_local_path = False @staticmethod @cache - def _bridge_mapping_types() -> tuple[Any, Any, Any, Any, Any, Any]: + def _bridge_mapping_types() -> tuple[Any, dict[Any, int]]: # Bridge is optional outside Megatron workers, so keep these imports local. from megatron.bridge.models.conversion.param_mapping import ( AutoMapping, @@ -79,14 +77,13 @@ def _bridge_mapping_types() -> tuple[Any, Any, Any, Any, Any, Any]: RowParallelMapping, ) - return ( - AutoMapping, - ColumnParallelMapping, - DirectMapping, - GatedMLPMapping, - ReplicatedMapping, - RowParallelMapping, - ) + return AutoMapping, { + ColumnParallelMapping: _COLUMN, + DirectMapping: _DIRECT, + GatedMLPMapping: _GATED, + ReplicatedMapping: _REPLICATED, + RowParallelMapping: _ROW, + } @staticmethod def _all_reduce_max(values: list[int]) -> list[int]: @@ -104,18 +101,10 @@ def _all_reduce_max(values: list[int]) -> list[int]: @classmethod def _local_mapping_kind(cls, task: Any) -> int: - ( - AutoMapping, - *mapping_types, - ) = cls._bridge_mapping_types() + AutoMapping, mapping_kinds = cls._bridge_mapping_types() mapping = task.mapping - for mapping_type, kind in zip( - mapping_types, - (_COLUMN, _DIRECT, _GATED, _REPLICATED, _ROW), - strict=True, - ): - if type(mapping) is mapping_type: - return kind + if kind := mapping_kinds.get(type(mapping)): + return kind if ( type(mapping) is AutoMapping and mapping.permute_dims is None @@ -156,24 +145,6 @@ def _is_padded_or_tied_weight(task: Any) -> bool: for name in hf_names ) - @staticmethod - def _can_project_task(task: Any, kind: int, *, identity_export: bool) -> bool: - if kind == _UNSUPPORTED or not identity_export: - return False - mapping = task.mapping - if getattr(mapping, "is_grouped_export", False) or getattr( - mapping, "is_adapter", False - ): - return False - if MegatronRemoteSparseRefit._is_padded_or_tied_weight(task): - return False - if kind == _GATED: - return isinstance(mapping.hf_param, dict) and set(mapping.hf_param) == { - "gate", - "up", - } - return isinstance(mapping.hf_param, str) - @staticmethod def _owns_policy_local_task(task: Any, *, replicated: bool = False) -> bool: if not torch.distributed.is_initialized(): @@ -198,17 +169,31 @@ def _owns_policy_local_task(task: Any, *, replicated: bool = False) -> bool: replica_count *= task.mapping.tp_size return replica_rank == sparse_name_shard(task.global_param_name, replica_count) + @classmethod + def _task_ownership(cls, task: Any, kind: int) -> tuple[torch.Tensor | None, bool]: + tensor = task.param_weight + replicated = kind in (_DIRECT, _REPLICATED) or ( + kind == _ROW and tensor is not None and tensor.ndim == 1 + ) + return ( + tensor + if tensor is not None + and cls._owns_policy_local_task(task, replicated=replicated) + else None, + replicated, + ) + def _policy_local_path_is_safe(self) -> bool: config = getattr(self._worker, "cfg", {}) - if config.get("quant_cfg") is not None: - return False ddp_config = config.get("megatron_cfg", {}).get( "distributed_data_parallel_config", {} ) - if ddp_config.get("use_custom_fsdp", False): - return False fp8_cfg = getattr(self._worker, "fp8_cfg", None) - return not fp8_cfg or not fp8_cfg.get("fp8_param", False) + return ( + config.get("quant_cfg") is None + and not ddp_config.get("use_custom_fsdp", False) + and (not fp8_cfg or not fp8_cfg.get("fp8_param", False)) + ) @staticmethod def _canonical_hf_name(task: Any, name: str) -> str: @@ -216,7 +201,7 @@ def _canonical_hf_name(task: Any, name: str) -> str: if not mapping.is_expert or mapping.ep_size == 1: return name - match = re.search(r"(\.experts\.)(\d+)(\.)", name) + match = re.search(r"(?<=\.experts\.)\d+(?=\.)", name) config = getattr(task.megatron_module, "config", None) num_experts = getattr(config, "num_moe_experts", None) if match is None or not isinstance(num_experts, int): @@ -227,9 +212,9 @@ def _canonical_hf_name(task: Any, name: str) -> str: f"{mapping.ep_size}." ) experts_per_rank = num_experts // mapping.ep_size - expert = int(match.group(2)) % experts_per_rank + expert = int(match.group()) % experts_per_rank expert += experts_per_rank * mapping.ep_rank - return f"{name[: match.start(2)]}{expert}{name[match.end(2) :]}" + return f"{name[: match.start()]}{expert}{name[match.end() :]}" @staticmethod def _projection( @@ -250,17 +235,28 @@ def _projection( @classmethod def _task_local_tensors( - cls, task: Any, kind: int - ) -> list[tuple[str, torch.Tensor, SparseShardProjection]]: - if task.param_weight is None: - return [] - + cls, task: Any, kind: int, *, identity_export: bool = True + ) -> list[tuple[str, torch.Tensor, SparseShardProjection]] | None: mapping = task.mapping - tensor = task.param_weight - replicated = kind in (_DIRECT, _REPLICATED) or ( - kind == _ROW and tensor.ndim == 1 - ) - if not cls._owns_policy_local_task(task, replicated=replicated): + hf_param = mapping.hf_param + if ( + kind == _UNSUPPORTED + or not identity_export + or getattr(mapping, "is_grouped_export", False) + or getattr(mapping, "is_adapter", False) + or cls._is_padded_or_tied_weight(task) + ): + return None + if kind == _GATED: + if not isinstance(hf_param, dict) or set(hf_param) != { + "gate", + "up", + }: + return None + elif not isinstance(hf_param, str): + return None + tensor, replicated = cls._task_ownership(task, kind) + if tensor is None: return [] if kind == _GATED: gate, up = torch.chunk(tensor, 2, dim=0) @@ -269,7 +265,9 @@ def _task_local_tensors( f"{task.global_param_name}:{role}", value, cls._projection( - cls._canonical_hf_name(task, str(mapping.hf_param[role])), + cls._canonical_hf_name( + task, str(cast(dict[str, Any], hf_param)[role]) + ), value, shard_dim=0, shard_rank=mapping.tp_rank, @@ -279,7 +277,7 @@ def _task_local_tensors( for role, value in (("gate", gate), ("up", up)) ] - name = cls._canonical_hf_name(task, str(mapping.hf_param)) + name = cls._canonical_hf_name(task, cast(str, hf_param)) projection = ( SparseShardProjection(name, tuple(tensor.shape), (0,) * tensor.ndim) if replicated @@ -308,14 +306,16 @@ def _prepare_paths(self) -> None: if not tasks or not self._policy_local_path_is_safe(): misc_tasks.extend(tasks) return - self._uses_policy_local_path = True kinds = self._all_reduce_max([self._local_mapping_kind(task) for task in tasks]) identity_export = self._bridge_exports_are_identity() self._filter_misc_tasks = identity_export projections = {} for task, kind in zip(tasks, kinds, strict=True): - if self._can_project_task(task, kind, identity_export=identity_export): - for key, tensor, projection in self._task_local_tensors(task, kind): + local_tensors = self._task_local_tensors( + task, kind, identity_export=identity_export + ) + if local_tensors is not None: + for key, tensor, projection in local_tensors: if key in projections: raise ValueError( f"Duplicate policy-local sparse shard {key!r}." @@ -326,76 +326,62 @@ def _prepare_paths(self) -> None: task_index = len(misc_tasks) misc_tasks.append(task) - if task.param_weight is None: - continue - replicated = kind in (_DIRECT, _REPLICATED) or ( - kind == _ROW and task.param_weight.ndim == 1 - ) - if not self._owns_policy_local_task(task, replicated=replicated): - continue - key = f"{task_index}:{task.global_param_name}" - self._misc_local_tensors.append((key, task.param_weight)) + tensor, _ = self._task_ownership(task, kind) + if tensor is not None: + key = f"{task_index}:{task.global_param_name}" + self._misc_local_tensors.append((key, tensor)) - if projections: - self._local_tracker = DeltaCompressionTracker( - self._delta_config, projections=projections - ) - if self._misc_local_tensors: - self._change_tracker = DeltaCompressionTracker(self._delta_config) + self._policy_tracker = DeltaCompressionTracker( + self._delta_config, projections=projections + ) def _iter_misc_params( self, conversion_tasks: list[Any] | None = None ) -> Iterable[tuple[str, torch.Tensor]]: - return self._worker._iter_params_with_optional_kv_scales( - conversion_tasks=( - self._misc_conversion_tasks - if conversion_tasks is None - else conversion_tasks - ) + tasks = ( + self._misc_conversion_tasks + if conversion_tasks is None + else conversion_tasks ) + return self._worker._iter_params_with_optional_kv_scales(conversion_tasks=tasks) def _changed_misc_tasks(self) -> tuple[list[Any], int, int]: assert self._misc_conversion_tasks is not None changed_keys: set[str] = set() changed = total = 0 - if self._change_tracker is not None: - changed_keys, changed, total = self._change_tracker.prepare_change_summary( + if self._misc_local_tensors: + assert self._policy_tracker is not None + changed_keys, changed, total = self._policy_tracker.prepare_change_summary( self._misc_local_tensors ) flags = [0] * len(self._misc_conversion_tasks) for key in changed_keys: flags[int(key.partition(":")[0])] = 1 flags = self._all_reduce_max(flags) - if any(flags) and not self._filter_misc_tasks: - flags = [1] * len(flags) - grouped_keys = { - task.mapping.group_key - for task, task_changed in zip( - self._misc_conversion_tasks, flags, strict=True - ) - if task_changed and getattr(task.mapping, "is_grouped_export", False) - } - flags = [ - task_changed - or ( - getattr(task.mapping, "is_grouped_export", False) - and task.mapping.group_key in grouped_keys - ) - for task, task_changed in zip( - self._misc_conversion_tasks, flags, strict=True - ) - ] - return ( - [ + if not any(flags): + tasks = [] + elif not self._filter_misc_tasks: + tasks = self._misc_conversion_tasks + else: + grouped_keys = { + task.mapping.group_key + for task, task_changed in zip( + self._misc_conversion_tasks, flags, strict=True + ) + if task_changed and getattr(task.mapping, "is_grouped_export", False) + } + tasks = [ task for task, task_changed in zip( self._misc_conversion_tasks, flags, strict=True ) if task_changed - ], - changed, - total, - ) + or ( + getattr(task.mapping, "is_grouped_export", False) + and task.mapping.group_key in grouped_keys + ) + ] + return tasks, changed, total def initialize_baseline( self, @@ -411,22 +397,20 @@ def initialize_baseline( shard_count=shard_count, transport=transport, ) - if not self._uses_policy_local_path: + policy_tracker = self._policy_tracker + if policy_tracker is None: snapshot(self._iter_misc_params(), delta_tracker=self._tracker) return self.refit_info() - def snapshot_local_baselines() -> None: - for tracker, tensors in ( - (self._local_tracker, self._local_tensors), - (self._change_tracker, self._misc_local_tensors), - ): - if tracker is not None: - snapshot(tensors, delta_tracker=tracker, partition="none") - with ThreadPoolExecutor( max_workers=1, thread_name_prefix="nrl-refit-policy-local" ) as executor: - local_future = executor.submit(snapshot_local_baselines) + local_future = executor.submit( + snapshot, + self._local_tensors + self._misc_local_tensors, + delta_tracker=policy_tracker, + partition="none", + ) snapshot( self._iter_misc_params(), delta_tracker=self._tracker, @@ -440,9 +424,9 @@ def refit_info(self) -> dict[str, tuple[tuple[int, ...], torch.dtype]]: name: (tuple(tensor.shape), tensor.dtype) for name, tensor in self._tracker.baseline.items() } - if self._local_tracker is not None: + if self._policy_tracker is not None: for name, tensor in self._local_tensors: - projection = self._local_tracker.projections[name] + projection = self._policy_tracker.projections[name] info[projection.name] = (projection.global_shape, tensor.dtype) return info @@ -470,7 +454,8 @@ def stream( shard_rank=shard_rank, shard_count=shard_count, ) - if not self._uses_policy_local_path: + policy_tracker = self._policy_tracker + if policy_tracker is None: result = send( self._iter_misc_params(), delta_tracker=self._tracker, @@ -481,11 +466,11 @@ def stream( max_workers=1, thread_name_prefix="nrl-refit-policy-local" ) as executor: local_future = None - if self._local_tracker is not None: + if self._local_tensors: local_future = executor.submit( send, self._local_tensors, - delta_tracker=self._local_tracker, + delta_tracker=policy_tracker, transfer_id=f"{transfer_id}-local", partition="none", ) @@ -511,7 +496,7 @@ def stream( return result def finish(self, succeeded: bool) -> None: - for tracker in (self._tracker, self._local_tracker, self._change_tracker): + for tracker in (self._tracker, self._policy_tracker): if tracker is None: continue if succeeded: diff --git a/nemo_rl/utils/weight_transfer_sparse_codec.py b/nemo_rl/utils/weight_transfer_sparse_codec.py index 667803452ed..6199a78eca1 100644 --- a/nemo_rl/utils/weight_transfer_sparse_codec.py +++ b/nemo_rl/utils/weight_transfer_sparse_codec.py @@ -124,24 +124,6 @@ def map_locations( return mapped -@dataclass(frozen=True) -class SparseSourceRoute: - """Canonical source view consumed by one vLLM worker.""" - - offset: int - strides: tuple[int, ...] - shape: tuple[int, ...] - linear: bool - - -@dataclass(frozen=True) -class SparseSourcePlan: - """Source views needed by one worker for a canonical HF tensor.""" - - routes: tuple[SparseSourceRoute, ...] = () - identity: bool = False - - def integer_dtype_for_element_size(element_size: int) -> torch.dtype: try: return _INTEGER_DTYPE_BY_SIZE[element_size] @@ -320,106 +302,6 @@ def iter_decoded_sparse_payload( ) -def map_sparse_locations( - locations: torch.Tensor, - source_offset: int, - source_strides: tuple[int, ...], - shape: tuple[int, ...], - linear: bool, - target_offset: int = 0, - target_strides: tuple[int, ...] | None = None, -) -> tuple[torch.Tensor, torch.Tensor]: - if linear: - end = source_offset + prod(shape) - keep = (locations >= source_offset) & (locations < end) - return locations + target_offset - source_offset, keep - mapped = torch.full_like(locations, target_offset) - reconstructed = torch.full_like(locations, source_offset) - relative = locations - source_offset - for size, source_stride, target_stride in zip( - shape, source_strides, target_strides or source_strides, strict=True - ): - coordinate = ( - torch.div(relative, source_stride, rounding_mode="floor").remainder(size) - if size > 1 - else torch.zeros_like(locations) - ) - reconstructed.add_(coordinate * source_stride) - mapped.add_(coordinate * target_stride) - return mapped, reconstructed == locations - - -def select_sparse_source_entries( - locations: torch.Tensor, - values: torch.Tensor, - plan: SparseSourcePlan, -) -> tuple[torch.Tensor, torch.Tensor]: - if plan.identity: - return locations, values - if not plan.routes: - return locations[:0], values[:0] - if all(route.linear for route in plan.routes): - ranges = sorted( - (route.offset, route.offset + prod(route.shape)) for route in plan.routes - ) - merged = [] - for start, end in ranges: - if merged and start <= merged[-1][1]: - merged[-1] = (merged[-1][0], max(merged[-1][1], end)) - else: - merged.append((start, end)) - bounds = torch.tensor( - [bound for interval in merged for bound in interval], - dtype=locations.dtype, - ) - offsets = torch.searchsorted(locations, bounds).tolist() - location_parts = [ - locations[start:end] for start, end in zip(offsets[::2], offsets[1::2]) - ] - value_parts = [ - values[start:end] for start, end in zip(offsets[::2], offsets[1::2]) - ] - if len(location_parts) == 1: - return location_parts[0], value_parts[0] - return torch.cat(location_parts), torch.cat(value_parts) - keep = torch.zeros(locations.shape, dtype=torch.bool) - for route in plan.routes: - _, route_keep = map_sparse_locations( - locations, route.offset, route.strides, route.shape, route.linear - ) - keep.logical_or_(route_keep) - return locations[keep], values[keep] - - -def partition_decoded_sparse_entries( - decoded: Iterable[DecodedSparseItem], - plans: Mapping[str, SparseSourcePlan], -) -> list[DecodedSparseItem]: - """Select the canonical entries consumed by one colocated worker.""" - selected_items = [] - for item, locations, values in decoded: - plan = plans[str(item["name"])] - selected_locations, selected_values = select_sparse_source_entries( - locations, values, plan - ) - if not selected_locations.numel(): - continue - selected_item = dict(item) - sample_location_tensor = torch.tensor( - item.get("verification_locations", ()), dtype=locations.dtype - ) - sample_value_tensor = torch.tensor( - item.get("verification_values", ()), dtype=values.dtype - ) - samples, sample_bits = select_sparse_source_entries( - sample_location_tensor, sample_value_tensor, plan - ) - selected_item["verification_locations"] = samples.tolist() - selected_item["verification_values"] = sample_bits.tolist() - selected_items.append((selected_item, selected_locations, selected_values)) - return selected_items - - def _encode_explicit_locations( locations: torch.Tensor, ) -> torch.Tensor: @@ -471,7 +353,6 @@ def prepare_sparse_delta_payload( ) -> PreparedTensorPayload: self._wait_for_baseline_commits() sparse_infos = [] - verification_sources = [] pending_updates = {} changed_elements = total_elements = 0 for name, tensor in tensors: @@ -509,14 +390,12 @@ def prepare_sparse_delta_payload( self.encoding, ) ) - if self.verification_samples: - verification_sources.append((payload_locations, values)) pending_updates[name] = (locations, current_values) with self._pending_updates_lock: self._pending_updates.update(pending_updates) payload = encode_sparse_infos(sparse_infos) - if verification_sources: - self._add_verification_samples(payload[2], verification_sources) + if self.verification_samples: + self._add_verification_samples(payload[2]) return payload, changed_elements, total_elements def prepare_change_summary( @@ -552,24 +431,22 @@ def _find_changes( def _add_verification_samples( self, metadata: list[dict[str, Any]], - sources: list[tuple[torch.Tensor, torch.Tensor]], ) -> None: - total = sum(int(locations.numel()) for locations, _ in sources) + sizes = [int(item["value_end"]) - int(item["value_start"]) for item in metadata] + total = sum(sizes) count = min(self.verification_samples, total) sample_ranks = [ ((2 * index + 1) * total) // (2 * count) for index in range(count) ] sample_index = offset = 0 - for item, (locations, values) in zip(metadata, sources, strict=True): - end = offset + locations.numel() + for item, size in zip(metadata, sizes, strict=True): + end = offset + size + samples = 0 while sample_index < count and sample_ranks[sample_index] < end: - local_index = sample_ranks[sample_index] - offset - location = int(locations[local_index]) - item.setdefault("verification_locations", []).append(location) - item.setdefault("verification_values", []).append( - int(values[local_index]) - ) + samples += 1 sample_index += 1 + if samples: + item["verification_samples"] = samples offset = end def on_sync_succeeded(self) -> None: diff --git a/tests/unit/models/generation/test_vllm_sparse_delta.py b/tests/unit/models/generation/test_vllm_sparse_delta.py index 4be5fa61aae..d7af593c934 100644 --- a/tests/unit/models/generation/test_vllm_sparse_delta.py +++ b/tests/unit/models/generation/test_vllm_sparse_delta.py @@ -12,8 +12,9 @@ # See the License for the specific language governing permissions and # limitations under the License. +import io import math -from types import MethodType, SimpleNamespace +from types import SimpleNamespace from typing import Any from unittest.mock import MagicMock @@ -22,22 +23,21 @@ from nemo_rl.models.generation.vllm.vllm_sparse_delta import ( VllmSparseDeltaApplier, - _SparseLoadTracer, ) from nemo_rl.utils.weight_transfer_sparse_codec import ( - SparseSourcePlan, - SparseSourceRoute, + SparseOperation, decode_sparse_tensor_payload_for_staging, encode_sparse_infos, integer_dtype_for_element_size, iter_decoded_sparse_payload, - partition_decoded_sparse_entries, ) class _NativeLoaderModel: def __init__(self, **targets: torch.Tensor) -> None: self.targets = targets + self.load_calls = 0 + self.loaded_names: list[str] = [] def parameters(self): return iter(self.targets.values()) @@ -46,7 +46,10 @@ def buffers(self): return iter(()) def load_weights(self, weights): + self.load_calls += 1 + loaded = set() for name, source in weights: + self.loaded_names.append(name) target = self.targets if name == "weight": target["identity"].copy_(source) @@ -70,7 +73,13 @@ def load_weights(self, weights): target["a"].copy_(-torch.exp(target["a"])) elif name == "transformed": target["identity"].copy_(source + 1) - return set() + elif name == "raise_after_copy": + target["identity"].copy_(source) + raise RuntimeError("loader failed") + else: + continue + loaded.add(name) + return loaded def _applier(model: Any) -> VllmSparseDeltaApplier: @@ -95,111 +104,50 @@ def _decode_staged(payload: Any) -> list[Any]: ) -def _apply_payload(applier: VllmSparseDeltaApplier, payload: Any) -> None: - decoded = _decode_staged(payload) - plans = applier.sparse_delta_source_plans([item for item, _, _ in decoded]) - applier._apply_decoded_sparse_weight_deltas( - partition_decoded_sparse_entries(decoded, plans) +def _payload( + name: str, + tensor: torch.Tensor, + locations: torch.Tensor | list[int], + values: torch.Tensor, + operation: SparseOperation = "overwrite", +) -> Any: + return encode_sparse_infos( + [(name, tensor, torch.as_tensor(locations), values, operation)] ) -def test_sparse_plan_prewarm_uses_native_loader_without_applying_values() -> None: +def _apply_payload(applier: VllmSparseDeltaApplier, payload: Any) -> None: + for item, locations, values in _decode_staged(payload): + applier._apply_decoded_item(item, locations, values) + + +def test_sparse_prewarm_reserves_largest_source_without_loading_weights() -> None: target = torch.zeros(8) applier = _applier(_NativeLoaderModel(identity=target)) applier.prewarm({"weight": ((8,), torch.float32)}) - assert applier._plan_cache["weight"].identity + assert applier._scratch.numel() == target.numel() * target.element_size() assert torch.equal(target, torch.zeros_like(target)) -def test_canonical_payload_is_partitioned_by_worker_source_plan() -> None: - tensor = torch.empty((8, 2), dtype=torch.float32) - locations = torch.tensor([1, 3, 8, 13]) - values = _bits(torch.tensor([1.0, 2.0, 3.0, 4.0])) - payload = encode_sparse_infos([("weight", tensor, locations, values, "overwrite")]) - payload[2][0].update( - verification_locations=[1, 13], - verification_values=[int(values[0]), int(values[3])], - ) - decoded = _decode_staged(payload) - - first = partition_decoded_sparse_entries( - decoded, - { - "weight": SparseSourcePlan( - routes=(SparseSourceRoute(0, (2, 1), (4, 2), True),) - ) - }, - ) - second = partition_decoded_sparse_entries( - decoded, - { - "weight": SparseSourcePlan( - routes=(SparseSourceRoute(8, (2, 1), (4, 2), True),) - ) - }, - ) - - first_item, first_locations, first_values = first[0] - second_item, second_locations, second_values = second[0] - assert first_locations.tolist() == [1, 3] - assert first_values.tolist() == values[:2].tolist() - assert first_item["verification_locations"] == [1] - assert second_locations.tolist() == [8, 13] - assert second_values.tolist() == values[2:].tolist() - assert second_item["verification_locations"] == [13] - - -def test_canonical_payload_partition_handles_strided_source_view() -> None: - values = torch.arange(8, dtype=torch.int32) - payload = encode_sparse_infos( - [ - ( - "weight", - torch.empty((8,), dtype=torch.float32), - torch.arange(8), - values, - "overwrite", - ) - ] - ) - partition = partition_decoded_sparse_entries( - _decode_staged(payload), - { - "weight": SparseSourcePlan( - routes=(SparseSourceRoute(1, (4, 1), (2, 2), False),) - ) - }, - ) - - _, locations, selected = partition[0] - assert locations.tolist() == [1, 2, 5, 6] - assert selected.tolist() == values[[1, 2, 5, 6]].tolist() +@pytest.mark.vllm +def test_sparse_prewarm_caches_rank_local_native_loader_skips() -> None: + target = torch.zeros(2) + model = _NativeLoaderModel(identity=target) + applier = _applier(model) + info = { + "weight": ((2,), torch.float32), + "skipped": ((2,), torch.float32), + } + applier.prewarm(info) + applier.discover_native_skips(info) + payload = _payload("skipped", target, [1], _bits(torch.tensor([3.0]))) + _apply_payload(applier, payload) -def test_canonical_payload_partitions_row_shards() -> None: - values = torch.arange(16, dtype=torch.int32) - payload = encode_sparse_infos( - [ - ( - "weight", - torch.empty((2, 8), dtype=torch.float32), - torch.arange(16), - values, - "overwrite", - ) - ] - ) - left = SparseSourcePlan(routes=(SparseSourceRoute(0, (8, 1), (2, 4), False),)) - right = SparseSourcePlan(routes=(SparseSourceRoute(4, (8, 1), (2, 4), False),)) - expected = {0: [0, 1, 2, 3, 8, 9, 10, 11], 1: [4, 5, 6, 7, 12, 13, 14, 15]} - for rank, plan in ((0, left), (1, right)): - _, locations, selected = partition_decoded_sparse_entries( - _decode_staged(payload), {"weight": plan} - )[0] - assert locations.tolist() == expected[rank] - assert selected.tolist() == values[expected[rank]].tolist() + assert model.loaded_names == ["weight", "skipped"] + assert torch.equal(target, torch.zeros_like(target)) @pytest.mark.vllm @@ -227,11 +175,14 @@ def test_backend_applies_decoded_sparse_payload_files() -> None: "first", "second" ) applier.prewarm.assert_called_once_with({"weight": ((8,), torch.float32)}) + applier.discover_native_skips.assert_called_once_with( + {"weight": ((8,), torch.float32)} + ) @pytest.mark.vllm def test_sparse_payload_batches_preserve_order(tmp_path) -> None: - applier = VllmSparseDeltaApplier(SimpleNamespace(), torch.device("cpu")) + applier = _applier(_NativeLoaderModel(identity=torch.zeros(1))) decoded_paths = [tmp_path / f"decoded-{index}.pt" for index in range(3)] decoded_payloads = [ ( @@ -244,6 +195,11 @@ def test_sparse_payload_batches_preserve_order(tmp_path) -> None: { "name": "weight", "index": index, + "shape": (1,), + "dtype": "float32", + "operation": "overwrite", + "index_encoding": "range", + "range_start": 0, "decoded_location_group": 0, "decoded_location_start": 0, "decoded_location_end": 1, @@ -258,10 +214,9 @@ def test_sparse_payload_batches_preserve_order(tmp_path) -> None: for path, payload in zip(decoded_paths, decoded_payloads, strict=True): torch.save(payload, path) decoded_applied: list[Any] = [] - applier.sparse_delta_source_plans = lambda _metadata: { - "weight": SparseSourcePlan(identity=True) - } - applier._apply_decoded_sparse_weight_deltas = decoded_applied.append + applier._apply_decoded_items = lambda items: decoded_applied.extend( + item for item, _, _ in items + ) result = applier.update_weights_from_decoded_sparse_payload( *(path.read_bytes() for path in decoded_paths) ) @@ -269,7 +224,7 @@ def test_sparse_payload_batches_preserve_order(tmp_path) -> None: *(str(path) for path in reversed(decoded_paths)) ) - assert [payload[0][0]["index"] for payload in decoded_applied] == [ + assert [item["index"] for item in decoded_applied] == [ 0, 1, 2, @@ -278,26 +233,36 @@ def test_sparse_payload_batches_preserve_order(tmp_path) -> None: 0, ] assert result["receiver_deserialize_s"] >= 0.0 - assert result["receiver_plan_s"] >= 0.0 + assert result["receiver_scratch_s"] >= 0.0 assert result["receiver_sparse_apply_s"] >= 0.0 assert decoded_result["receiver_deserialize_s"] >= 0.0 - assert decoded_result["receiver_partition_s"] >= 0.0 + + +@pytest.mark.vllm +def test_sparse_payload_batch_uses_one_streaming_native_loader_call() -> None: + identity = torch.zeros(2) + scale = torch.zeros(2) + model = _NativeLoaderModel(identity=identity, scale=scale) + serialized = [] + for payload in ( + _payload("weight", identity, [0], _bits(torch.tensor([2.0]))), + _payload("weight_scale_inv", scale, [1], _bits(torch.tensor([3.0]))), + ): + buffer = io.BytesIO() + torch.save(decode_sparse_tensor_payload_for_staging(payload), buffer) + serialized.append(buffer.getvalue()) + + _applier(model).update_weights_from_decoded_sparse_payload(*serialized) + + assert model.load_calls == 1 + assert torch.equal(identity, torch.tensor([2.0, 0.0])) + assert torch.equal(scale, torch.tensor([0.0, 3.0])) @pytest.mark.vllm def test_decoded_sparse_payload_converts_compact_locations_for_apply(tmp_path) -> None: target = torch.zeros(8) - payload = encode_sparse_infos( - [ - ( - "weight", - target, - torch.tensor([1, 5]), - _bits(torch.tensor([2.0, 6.0])), - "overwrite", - ) - ] - ) + payload = _payload("weight", target, [1, 5], _bits(torch.tensor([2.0, 6.0]))) decoded = decode_sparse_tensor_payload_for_staging(payload) assert decoded[0][0].dtype == torch.int32 path = tmp_path / "decoded.pt" @@ -308,11 +273,11 @@ def test_decoded_sparse_payload_converts_compact_locations_for_apply(tmp_path) - ).update_weights_from_decoded_sparse_payload_files(str(path)) assert torch.equal(target, torch.tensor([0.0, 2.0, 0.0, 0.0, 0.0, 6.0, 0.0, 0.0])) - assert result["receiver_partition_s"] >= 0.0 + assert result["receiver_sparse_apply_s"] >= 0.0 @pytest.mark.vllm -def test_native_loaders_compile_sparse_placement() -> None: +def test_native_loaders_apply_sparse_overwrite_without_family_plans() -> None: targets = { "identity": torch.zeros(4), "qkv": torch.zeros(8, 2), @@ -341,7 +306,6 @@ def test_native_loaders_compile_sparse_placement() -> None: [0, 1], [math.log(3.0), math.log(2.0)], ), - ("model.layers.0.mlp.experts.7.gate_proj.weight", (8, 2), [8], [9]), ] payload = encode_sparse_infos( [ @@ -355,20 +319,10 @@ def test_native_loaders_compile_sparse_placement() -> None: for name, shape, locations, values in infos ] ) - payload[2][-1].update( - verification_locations=[8], - verification_values=[int(_bits(torch.tensor([9.0]))[0])], - ) - payload[2][7].update( - verification_locations=[0, 1], - verification_values=[ - int(value) for value in _bits(torch.tensor([math.log(3.0), math.log(2.0)])) - ], - ) + payload[2][7]["verification_samples"] = 2 applier = _applier(_NativeLoaderModel(**targets)) _apply_payload(applier, payload) - plans = applier.sparse_delta_source_plans(payload[2]) verification = applier.finish_sparse_delta_refit() assert torch.equal(targets["identity"], torch.tensor([0.0, 1.0, 2.0, 0.0])) @@ -378,9 +332,6 @@ def test_native_loaders_compile_sparse_placement() -> None: assert targets["w2"].view(-1)[[8, 11, 12, 15]].tolist() == [5, 5, 5, 5] assert targets["mamba"].view(-1)[[0, 3, 4, 11]].tolist() == [6, 6, 6, 6] assert torch.allclose(targets["a"], torch.tensor([-3.0, -2.0])) - assert plans["model.layers.0.self_attn.k_proj.weight"].routes == ( - SparseSourceRoute(4, (2, 1), (2, 2), True), - ) assert verification["verification_candidates"] == 2 assert verification["verification_samples"] == 2 assert verification["verification_exact_mismatches"] == 0 @@ -450,129 +401,56 @@ def test_xor_applies_through_packed_native_loaders() -> None: @pytest.mark.vllm -def test_vllm_native_loader_geometry() -> None: - from vllm.model_executor.layers.fused_moe.layer import FusedMoE - from vllm.model_executor.layers.linear import ( - MergedColumnParallelLinear, - QKVParallelLinear, - ) - from vllm.model_executor.layers.mamba.mamba_mixer2 import ( - mamba_v2_sharded_weight_loader, - ) +def test_native_loader_explicit_skip_is_accepted() -> None: + model = _NativeLoaderModel(identity=torch.zeros(1)) + payload = _payload("skipped", torch.empty(1), [0], _bits(torch.tensor([1.0]))) + item, locations, values = _decode_staged(payload)[0] - def trace(target, source_shape, load): - source = torch.empty(source_shape, device="meta") - tracer = _SparseLoadTracer([target], {"source": source}) - with tracer: - load(source) - return tracer.copies["source"] - - qkv = torch.nn.Parameter(torch.zeros(8, 2)) - qkv.output_dim = 0 - qkv_layer = SimpleNamespace( - num_heads=4, - num_kv_heads=2, - num_kv_head_replicas=2, - head_size=1, - v_head_size=1, - tp_rank=3, - ) - qkv_layer.validate_shard_id = MethodType( - QKVParallelLinear.validate_shard_id, qkv_layer - ) - qkv_copies = trace( - qkv, - (4, 2), - lambda source: QKVParallelLinear.weight_loader(qkv_layer, qkv, source, "k"), - ) + _applier(model)._apply_decoded_item(item, locations, values) - merged = torch.nn.Parameter(torch.zeros(8, 2)) - merged.output_dim = 0 - merged_layer = SimpleNamespace(output_sizes=[8, 8], tp_size=2, tp_rank=1) - merged_layer.validate_shard_id = MethodType( - MergedColumnParallelLinear.validate_shard_id, merged_layer - ) - merged_copies = trace( - merged, - (8, 2), - lambda source: MergedColumnParallelLinear.weight_loader( - merged_layer, merged, source, 1 - ), - ) + assert torch.equal(model.targets["identity"], torch.zeros(1)) - expert = torch.nn.Parameter(torch.zeros(2, 8, 2)) - moe = SimpleNamespace( - moe_config=SimpleNamespace(is_act_and_mul=True), - _get_hidden_dim=FusedMoE._get_hidden_dim, - _narrow_expert_data_for_padding=FusedMoE._narrow_expert_data_for_padding, - ) - expert_copies = trace( - expert, - (8, 2), - lambda source: FusedMoE._load_w13(moe, expert.data[1], 0, "w3", source, 1), - ) - mamba = torch.nn.Parameter(torch.zeros(10, 1, 2)) - mamba_loader = mamba_v2_sharded_weight_loader( - [(8, 0, False), (4, 2, True), (4, 2, True), (4, 0, False)], 2, 1 - ) - mamba_copies = trace(mamba, (16, 1, 2), lambda source: mamba_loader(mamba, source)) +@pytest.mark.vllm +def test_native_loader_claim_without_copy_fails_closed() -> None: + model = _NativeLoaderModel(identity=torch.zeros(1)) + model.load_weights = lambda weights: {name for name, _ in weights} + payload = _payload("weight", torch.empty(1), [0], _bits(torch.tensor([1.0]))) - cases = ( - (qkv_copies, [0, 4, 7], [8, 11]), - (merged_copies, [0, 8, 15], [8, 15]), - (expert_copies, [0, 8, 15], [24, 31]), - ( - mamba_copies, - [0, 8, 15, 16, 19, 20, 23, 28, 31], - [0, 7, 8, 11, 12, 15, 16, 19], - ), + with pytest.raises(RuntimeError, match="without a supported target copy"): + _apply_payload(_applier(model), payload) + + +@pytest.mark.vllm +def test_sparse_overwrite_preserves_unselected_transform_inputs() -> None: + target = torch.tensor([-2.0, -4.0, -6.0, -8.0]) + payload = _payload( + "backbone.layers.0.mixer.A_log", + target, + [1], + _bits(torch.tensor([math.log(3.0)])), ) - for copies, source_locations, expected_rows in cases: - mapped = [ - VllmSparseDeltaApplier._map_copy( - torch.tensor(source_locations), - torch.ones(len(source_locations)), - copy, - )[0] - for copy in copies - ] - assert torch.cat(mapped).tolist() == expected_rows + + _apply_payload(_applier(_NativeLoaderModel(a=target)), payload) + + assert torch.allclose(target, torch.tensor([-2.0, -3.0, -6.0, -8.0])) @pytest.mark.vllm -def test_unknown_native_loader_fails_closed() -> None: - model = _NativeLoaderModel(identity=torch.zeros(1)) - for name, error in (("unknown", "did not place"), ("transformed", "transform")): - payload = encode_sparse_infos( - [ - ( - name, - torch.empty(1), - torch.tensor([0]), - _bits(torch.tensor([1.0])), - "overwrite", - ) - ], - ) - with pytest.raises(RuntimeError, match=error): - _apply_payload(_applier(model), payload) +def test_sparse_overwrite_rolls_back_loader_failure() -> None: + target = torch.tensor([1.0, 2.0]) + payload = _payload("raise_after_copy", target, [1], _bits(torch.tensor([3.0]))) + + with pytest.raises(RuntimeError, match="loader failed"): + _apply_payload(_applier(_NativeLoaderModel(identity=target)), payload) + + assert torch.equal(target, torch.tensor([1.0, 2.0])) @pytest.mark.vllm def test_unknown_sparse_operation_fails_closed() -> None: target = torch.zeros(1) - payload = encode_sparse_infos( - [ - ( - "weight", - target, - torch.tensor([0]), - _bits(torch.tensor([1.0])), - "overwrite", - ) - ] - ) + payload = _payload("weight", target, [0], _bits(torch.tensor([1.0]))) payload[2][0]["operation"] = "unknown" with pytest.raises(ValueError, match="Unsupported sparse-refit operation"): @@ -592,32 +470,18 @@ def test_sparse_delta_verification_compares_replacement( ) -> None: target = torch.tensor([1.0, initial, 3.0, initial]) replacement = torch.tensor([initial + 4.0, initial + 4.0]) - payload = encode_sparse_infos( - [ - ( - "weight", - target, - torch.tensor([1, 3]), - _bits(replacement), - "overwrite", - ) - ], - ) - payload[2][0].update( - verification_locations=[1, 3], - verification_values=[ - int(value) - for value in _bits( - torch.tensor([initial + verified_value, initial + verified_value]) - ) - ], - ) + payload = _payload("weight", target, [1, 3], _bits(replacement)) + payload[2][0]["verification_samples"] = 2 applier = _applier(_NativeLoaderModel(identity=target)) _apply_payload(applier, payload) + target[[1, 3]] = initial + verified_value result = applier.finish_sparse_delta_refit() - assert torch.equal(target, torch.tensor([1.0, initial + 4.0, 3.0, initial + 4.0])) + assert torch.equal( + target, + torch.tensor([1.0, initial + verified_value, 3.0, initial + verified_value]), + ) assert result["verification_candidates"] == 2 assert result["verification_samples"] == 2 assert result["verification_exact_mismatches"] == 2 * exact_mismatches @@ -653,13 +517,8 @@ def test_fp8_weight_and_scale_use_exact_bit_overwrite() -> None: ), ] ) - payload[2][0].update( - verification_locations=[1, 2], verification_values=[0x41, 0x7F] - ) - payload[2][1].update( - verification_locations=[0], - verification_values=[int(current_scale.view(torch.int32)[0])], - ) + payload[2][0]["verification_samples"] = 2 + payload[2][1]["verification_samples"] = 1 applier = _applier(_NativeLoaderModel(identity=target, scale=scale)) _apply_payload(applier, payload) @@ -682,10 +541,7 @@ def test_xor_applies_exact_bits_and_replay_reverts() -> None: locations = torch.tensor([1, 2]) xor_values = _bits(current)[locations].bitwise_xor(_bits(baseline)[locations]) payload = encode_sparse_infos([("weight", current, locations, xor_values, "xor")]) - payload[2][0].update( - verification_locations=locations.tolist(), - verification_values=[int(value) for value in xor_values], - ) + payload[2][0]["verification_samples"] = 2 applier = _applier(_NativeLoaderModel(identity=target)) _apply_payload(applier, payload) @@ -703,17 +559,7 @@ def test_xor_applies_exact_bits_and_replay_reverts() -> None: def test_overwrite_casts_absolute_source_values() -> None: target = torch.zeros(2, dtype=torch.float16) source = torch.tensor([1.25, -2.5], dtype=torch.float32) - payload = encode_sparse_infos( - [ - ( - "weight", - source, - torch.tensor([0, 1]), - _bits(source), - "overwrite", - ) - ] - ) + payload = _payload("weight", source, [0, 1], _bits(source)) _apply_payload(_applier(_NativeLoaderModel(identity=target)), payload) @@ -728,13 +574,13 @@ def test_overwrite_casts_absolute_source_values() -> None: "backbone.layers.0.mixer.A_log", torch.tensor([math.log(2.0)]), {"a": torch.tensor([-1.0])}, - "transformed weight", + "transforms its input", ), ( "weight", torch.tensor([1.0], dtype=torch.float32), {"identity": torch.zeros(1, dtype=torch.float16)}, - "dtypes differ", + "dtypes must match", ), ], ) @@ -744,16 +590,14 @@ def test_xor_rejects_non_bitwise_compatible_targets( targets: dict[str, torch.Tensor], error: str, ) -> None: - payload = encode_sparse_infos( - [(name, source, torch.tensor([0]), _bits(source), "xor")] - ) + payload = _payload(name, source, [0], _bits(source), "xor") with pytest.raises(RuntimeError, match=error): _apply_payload(_applier(_NativeLoaderModel(**targets)), payload) @pytest.mark.vllm -def test_xor_rejects_overlapping_target_mappings() -> None: +def test_xor_rejects_overlapping_native_loader_copies() -> None: class RepeatedCopyModel(torch.nn.Module): def __init__(self) -> None: super().__init__() @@ -765,17 +609,13 @@ def load_weights(self, weights) -> None: self.weight.copy_(source) source = torch.tensor([1.0, 2.0]) - payload = encode_sparse_infos( - [ - ( - "weight", - source, - torch.tensor([0, 1]), - _bits(source).bitwise_xor(_bits(torch.zeros_like(source))), - "xor", - ) - ] + payload = _payload( + "weight", + source, + [0, 1], + _bits(source).bitwise_xor(_bits(torch.zeros_like(source))), + "xor", ) - with pytest.raises(RuntimeError, match="target mappings overlap"): + with pytest.raises(RuntimeError, match="overlapping target copies"): _apply_payload(_applier(RepeatedCopyModel()), payload) diff --git a/tests/unit/models/generation/test_vllm_sparse_refit.py b/tests/unit/models/generation/test_vllm_sparse_refit.py index 259f859c9da..8e906bf4e01 100644 --- a/tests/unit/models/generation/test_vllm_sparse_refit.py +++ b/tests/unit/models/generation/test_vllm_sparse_refit.py @@ -91,10 +91,7 @@ def _serialized_sparse_payload() -> bytes: ) ] ) - payload[2][0].update( - verification_locations=[1, 7], - verification_values=[1, 4], - ) + payload[2][0]["verification_samples"] = 2 buffer = io.BytesIO() torch.save(payload, buffer) return buffer.getvalue() diff --git a/tests/unit/models/policy/test_megatron_remote_sparse_refit.py b/tests/unit/models/policy/test_megatron_remote_sparse_refit.py index f6fa839a3a0..d3130d4f4cf 100644 --- a/tests/unit/models/policy/test_megatron_remote_sparse_refit.py +++ b/tests/unit/models/policy/test_megatron_remote_sparse_refit.py @@ -71,11 +71,13 @@ def _install_mapping_types(monkeypatch, remote_refit_type): staticmethod( lambda: ( _AutoMapping, - _ColumnMapping, - _DirectMapping, - _GatedMapping, - _ReplicatedMapping, - _RowMapping, + { + _ColumnMapping: megatron_remote_sparse_refit._COLUMN, + _DirectMapping: megatron_remote_sparse_refit._DIRECT, + _GatedMapping: megatron_remote_sparse_refit._GATED, + _ReplicatedMapping: megatron_remote_sparse_refit._REPLICATED, + _RowMapping: megatron_remote_sparse_refit._ROW, + }, ) ), ) @@ -86,19 +88,29 @@ def _install_mapping_types(monkeypatch, remote_refit_type): ) -def test_remote_sparse_stream_drains_cuda_before_return(monkeypatch): - class Worker: - cfg = {} - fp8_cfg = None - model = object() - megatron_bridge = SimpleNamespace(get_conversion_tasks=lambda _models: []) - - @staticmethod - def _iter_params_with_optional_kv_scales(*, conversion_tasks=None): - assert conversion_tasks == [] +def _worker(tasks=(), *, fp8_cfg=None, export=None): + if export is None: + + def export(*, conversion_tasks=None): return iter(()) - worker = Worker() + return SimpleNamespace( + cfg={}, + fp8_cfg=fp8_cfg, + model=object(), + megatron_bridge=SimpleNamespace( + get_conversion_tasks=lambda _models: list(tasks) + ), + _iter_params_with_optional_kv_scales=export, + ) + + +def test_remote_sparse_stream_drains_cuda_before_return(monkeypatch): + def export(*, conversion_tasks=None): + assert conversion_tasks == [] + return iter(()) + + worker = _worker(export=export) remote_refit = MegatronRemoteSparseRefit(worker, _DELTA_CONFIG) result = {"payloads": 1, "changed_elements": 2, "total_elements": 3} events = [] @@ -132,17 +144,10 @@ def stream(*_args, **_kwargs): def test_remote_sparse_stream_combines_local_and_misc_paths(monkeypatch): - class Worker: - @staticmethod - def _iter_params_with_optional_kv_scales(*, conversion_tasks=None): - assert conversion_tasks == [] - return iter(()) - - remote_refit = MegatronRemoteSparseRefit(Worker(), _DELTA_CONFIG) - remote_refit._local_tracker = object() + remote_refit = MegatronRemoteSparseRefit(_worker(), _DELTA_CONFIG) + remote_refit._policy_tracker = object() remote_refit._local_tensors = [("local", torch.ones(1))] remote_refit._misc_conversion_tasks = [] - remote_refit._uses_policy_local_path = True monkeypatch.setattr(remote_refit, "_changed_misc_tasks", lambda: ([], 5, 6)) calls = [] @@ -288,28 +293,22 @@ class _TransformedMapping(_AutoMapping): global_param_name="decoder.linear_qkv.weight", ) - class Worker: - cfg = {} - fp8_cfg = None - model = object() - megatron_bridge = SimpleNamespace(get_conversion_tasks=lambda _models: [task]) - _install_mapping_types(monkeypatch, MegatronRemoteSparseRefit) monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") - remote_refit = MegatronRemoteSparseRefit(Worker(), _DELTA_CONFIG) + remote_refit = MegatronRemoteSparseRefit(_worker([task]), _DELTA_CONFIG) remote_refit._prepare_paths() assert remote_refit._misc_conversion_tasks == [task] assert remote_refit._misc_local_tensors == [("0:decoder.linear_qkv.weight", tensor)] - assert remote_refit._change_tracker is not None - remote_refit._change_tracker.snapshot_baseline(remote_refit._misc_local_tensors) + assert remote_refit._policy_tracker is not None + remote_refit._policy_tracker.snapshot_baseline(remote_refit._misc_local_tensors) tensor[1] = 5 changed_tasks, changed, total = remote_refit._changed_misc_tasks() assert changed_tasks == [task] assert (changed, total) == (1, 3) - remote_refit._change_tracker.on_sync_succeeded() + remote_refit._policy_tracker.on_sync_succeeded() assert remote_refit._changed_misc_tasks() == ([], 0, 3) @@ -334,24 +333,18 @@ class _TransformedMapping(_AutoMapping): ), ] - class Worker: - cfg = {} - fp8_cfg = None - model = object() - megatron_bridge = SimpleNamespace(get_conversion_tasks=lambda _models: tasks) - _install_mapping_types(monkeypatch, MegatronRemoteSparseRefit) monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") - remote_refit = MegatronRemoteSparseRefit(Worker(), _XOR_CONFIG) + remote_refit = MegatronRemoteSparseRefit(_worker(tasks), _XOR_CONFIG) remote_refit._prepare_paths() - assert remote_refit._local_tracker is not None - assert remote_refit._local_tracker.encoding == "xor" + assert remote_refit._policy_tracker is not None + assert remote_refit._policy_tracker.encoding == "xor" assert remote_refit._tracker.encoding == "overwrite" - remote_refit._local_tracker.snapshot_baseline(remote_refit._local_tensors) + remote_refit._policy_tracker.snapshot_baseline(remote_refit._local_tensors) direct[0, 0] = 5 - direct_metadata = remote_refit._local_tracker.prepare_sparse_delta_payload( + direct_metadata = remote_refit._policy_tracker.prepare_sparse_delta_payload( remote_refit._local_tensors )[0][2] assert [item["operation"] for item in direct_metadata] == ["xor"] @@ -452,18 +445,13 @@ def test_remote_sparse_balances_tasks_across_equivalent_replicas(monkeypatch): def test_remote_sparse_fp8_policy_keeps_full_export_path(monkeypatch): task = object() - class Worker: - cfg = {} - fp8_cfg = {"fp8_param": True} - model = object() - megatron_bridge = SimpleNamespace(get_conversion_tasks=lambda _models: [task]) + def export(*, conversion_tasks): + assert conversion_tasks == [task] + return iter(()) - @staticmethod - def _iter_params_with_optional_kv_scales(*, conversion_tasks): - assert conversion_tasks == [task] - return iter(()) - - remote_refit = MegatronRemoteSparseRefit(Worker(), _DELTA_CONFIG) + remote_refit = MegatronRemoteSparseRefit( + _worker([task], fp8_cfg={"fp8_param": True}, export=export), _DELTA_CONFIG + ) snapshots = [] monkeypatch.setattr( "nemo_rl.models.policy.workers.megatron_remote_sparse_refit." @@ -475,4 +463,4 @@ def _iter_params_with_optional_kv_scales(*, conversion_tasks): assert snapshots == [[]] assert remote_refit._misc_conversion_tasks == [task] - assert not remote_refit._uses_policy_local_path + assert remote_refit._policy_tracker is None diff --git a/tests/unit/utils/test_weight_transfer_remote_sparse.py b/tests/unit/utils/test_weight_transfer_remote_sparse.py index f01f2da670a..ecbe361c3eb 100644 --- a/tests/unit/utils/test_weight_transfer_remote_sparse.py +++ b/tests/unit/utils/test_weight_transfer_remote_sparse.py @@ -65,6 +65,16 @@ def prepare_sparse_delta_payload(chunk): return payload, count, count +class _BaselineNamesTracker: + sparse_bucket_size_bytes = 4 + + def __init__(self) -> None: + self.names = [] + + def snapshot_baseline(self, chunk) -> None: + self.names.extend(name for name, _tensor in chunk) + + def _stream_sparse_test_payloads(tensors, send_payload): return weight_transfer_remote_sparse.stream_sparse_delta_payloads( tensors, @@ -113,7 +123,7 @@ def test_delta_tracker_change_summary_is_transactional(monkeypatch) -> None: assert tracker.prepare_change_summary(tensors) == (set(), 0, 4) -def test_delta_tracker_emits_bounded_delta_samples(monkeypatch) -> None: +def test_delta_tracker_emits_bounded_verification_budget(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") monkeypatch.setenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "2") tracker = _delta_tracker() @@ -125,11 +135,7 @@ def test_delta_tracker_emits_bounded_delta_samples(monkeypatch) -> None: [("weight", tensor)] ) - assert metadata[0]["verification_locations"] == [1, 3] - assert metadata[0]["verification_values"] == [ - int(tensor.view(torch.int32)[1]), - int(tensor.view(torch.int32)[3]), - ] + assert metadata[0]["verification_samples"] == 2 assert metadata[0]["operation"] == "overwrite" assert (changed, total) == (2, 4) @@ -205,7 +211,7 @@ def test_delta_tracker_projects_local_shards_to_hf_locations( assert (changed, total) == (2, 4) assert metadata[0]["name"] == "hf.weight" assert metadata[0]["shape"] == projection.global_shape - assert metadata[0]["verification_locations"] == expected_locations + assert metadata[0]["verification_samples"] == 2 assert ( sparse_locations_for_item(metadata[0], locations, device="cpu").tolist() == expected_locations @@ -253,8 +259,6 @@ def test_delta_tracker_encodes_fp8_weight_and_scale_bits( torch.float8_e4m3fn ) scale = torch.tensor([1.0, 2.0], dtype=torch.float32) - baseline_weight = weight.clone() - baseline_scale = scale.clone() tracker.snapshot_baseline([("weight", weight), ("weight_scale_inv", scale)]) weight.view(torch.uint8)[1] = 0x41 scale[0] = 1.5 @@ -269,15 +273,7 @@ def test_delta_tracker_encodes_fp8_weight_and_scale_bits( assert len(value_groups) == 2 assert [item["operation"] for item in metadata] == [encoding, encoding] assert [item["dtype"] for item in metadata] == ["float8_e4m3fn", "float32"] - expected_weight = int(weight.view(torch.uint8)[1]) - expected_scale = int(scale.view(torch.int32)[0]) - if encoding == "xor": - expected_weight ^= int(baseline_weight.view(torch.uint8)[1]) - expected_scale ^= int(baseline_scale.view(torch.int32)[0]) - assert [item["verification_values"] for item in metadata] == [ - [expected_weight], - [expected_scale], - ] + assert [item["verification_samples"] for item in metadata] == [1, 1] assert [ sparse_locations_for_item(item, locations, device="cpu").tolist() for item in metadata @@ -457,17 +453,7 @@ def test_sparse_baseline_snapshots_only_owned_export_chunks( monkeypatch, capsys ) -> None: monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "4") - - class Tracker: - sparse_bucket_size_bytes = 4 - - def __init__(self) -> None: - self.names = [] - - def snapshot_baseline(self, chunk) -> None: - self.names.extend(name for name, _tensor in chunk) - - tracker = Tracker() + tracker = _BaselineNamesTracker() weight_transfer_remote_sparse.init_sparse_delta_baseline_from_iterator( [(f"weight-{index}", torch.tensor([float(index)])) for index in range(4)], delta_tracker=tracker, @@ -479,7 +465,7 @@ def snapshot_baseline(self, chunk) -> None: assert tracker.names == ["weight-1", "weight-3"] assert "chunks=4" in capsys.readouterr().out - tracker = Tracker() + tracker = _BaselineNamesTracker() weight_transfer_remote_sparse.init_sparse_delta_baseline_from_iterator( [(f"weight-{index}", torch.tensor([float(index)])) for index in range(4)], delta_tracker=tracker, @@ -494,19 +480,10 @@ def snapshot_baseline(self, chunk) -> None: def test_sparse_name_partition_is_stable_for_filtered_exports(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "4") - class Tracker: - sparse_bucket_size_bytes = 4 - - def __init__(self) -> None: - self.names = [] - - def snapshot_baseline(self, chunk) -> None: - self.names.extend(name for name, _tensor in chunk) - tensors = [(f"weight-{index}", torch.tensor([float(index)])) for index in range(8)] owners = [] for rank in range(2): - tracker = Tracker() + tracker = _BaselineNamesTracker() weight_transfer_remote_sparse.init_sparse_delta_baseline_from_iterator( tensors, delta_tracker=tracker, @@ -522,7 +499,7 @@ def snapshot_baseline(self, chunk) -> None: filtered = tensors[::2] for rank in range(2): - tracker = Tracker() + tracker = _BaselineNamesTracker() weight_transfer_remote_sparse.init_sparse_delta_baseline_from_iterator( filtered, delta_tracker=tracker, From 69358872477459e1e5734cea09f33a3048871974 Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Tue, 14 Jul 2026 15:10:21 -0700 Subject: [PATCH 11/18] Clean up Signed-off-by: Hollow Man --- docs/design-docs/sparse-delta-refit.md | 431 +++++++++--------- nemo_rl/algorithms/grpo.py | 31 +- .../models/generation/vllm/vllm_backend.py | 12 +- .../generation/vllm/vllm_sparse_delta.py | 313 +++++-------- .../generation/vllm/vllm_sparse_refit.py | 5 +- .../policy/workers/megatron_policy_worker.py | 4 +- .../workers/megatron_remote_sparse_refit.py | 33 +- .../utils/weight_transfer_remote_sparse.py | 123 ++--- nemo_rl/utils/weight_transfer_sparse_codec.py | 101 ++-- nemo_rl/utils/weight_transfer_zmq.py | 28 +- .../vllm_remote_sparse_weight_synchronizer.py | 54 ++- .../generation/test_vllm_sparse_delta.py | 238 ++++------ .../generation/test_vllm_sparse_refit.py | 14 +- .../models/megatron/test_community_import.py | 10 +- .../test_megatron_remote_sparse_refit.py | 112 ++--- .../test_weight_transfer_remote_sparse.py | 137 ++---- ..._vllm_remote_sparse_weight_synchronizer.py | 77 ++-- tools/refit_bandwidth_calculator.py | 38 +- 18 files changed, 723 insertions(+), 1038 deletions(-) diff --git a/docs/design-docs/sparse-delta-refit.md b/docs/design-docs/sparse-delta-refit.md index 21a0e3ac038..b84ac451b4b 100644 --- a/docs/design-docs/sparse-delta-refit.md +++ b/docs/design-docs/sparse-delta-refit.md @@ -1,16 +1,14 @@ # Remote Sparse-Delta vLLM Refit -Remote sparse-delta refit updates non-colocated vLLM workers without sending a -full checkpoint after every optimizer step. Megatron workers compare every -uniquely owned MCore tensor against a policy-local CPU baseline. Exact affine -mappings emit sparse Hugging Face (HF) coordinates directly; only changed tasks -whose conversion is not affine traverse Megatron Bridge. S3 or ZeroMQ carries -the resulting payloads, and each native vLLM weight loader applies its canonical -HF update to the local TP or EP destination. - -The feature is opt-in. Its synchronizer, codec, transports, receiver queue, and -native-loader apply engine are separate from existing NCCL, CUDA IPC, and -packed refit paths. +Remote sparse refit updates non-colocated vLLM workers without transferring a +full checkpoint after every optimizer step. Megatron workers compare uniquely +owned MCore shards against a policy-local CPU baseline. Affine Bridge mappings +emit changed values in canonical Hugging Face (HF) coordinates; only residual +conversion tasks use `export_hf_weights()`. S3 and ZeroMQ share the codec, +pipeline, receiver, native-loader apply path, and commit protocol. + +The feature is opt-in and does not change existing NCCL, CUDA IPC, or packed +refit behavior. ## Supported scope @@ -18,15 +16,17 @@ Remote sparse refit requires: - a non-colocated Megatron policy and vLLM generation backend; - the same initial HF checkpoint on both clusters; -- BF16 or FP16, unquantized rollout weights; +- BF16 or FP16 unquantized rollout weights; - `kv_cache_dtype: auto`; and - a `delta_compression` configuration. -Configuration validation rejects FP8 weights, FP8 KV-cache scales, -`quant_cfg`, `real_quant`, colocated inference, and non-Megatron policies. +Validation rejects `quant_cfg`, `real_quant`, colocated or non-Megatron +deployments, FP8 rollout weights, and FP8 KV-cache scales. The codec and +overwrite apply path retain FP8 bit patterns, but they do not generate block +scales or KV-cache scales, so that alone is not end-to-end FP8 support. Synchronous and asynchronous vLLM engines are supported, but the weight-version -transition is synchronous: generation pauses until all payloads are applied and -the global flush completes. +transition remains synchronous: generation pauses until every payload is +applied and the global flush completes. ## Architecture @@ -51,7 +51,7 @@ flowchart LR subgraph G["vLLM generation cluster"] H["HTTP receiver"] Q["Eager node staging and bounded FIFO apply queue"] - A["Dense scratch and native-loader apply"] + A["Reusable dense scratch and native load_weights()"] H --> Q --> A end @@ -65,41 +65,58 @@ apply engine, and commit protocol.* | Responsibility | Implementation | |---|---| | Coordinate one transfer and commit | [`vllm_remote_sparse_weight_synchronizer.py`](../../nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py) | -| Adapt Megatron workers | [`megatron_remote_sparse_refit.py`](../../nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py) | +| Adapt Megatron workers and assign ownership | [`megatron_remote_sparse_refit.py`](../../nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py) | | Track baselines and encode deltas | [`weight_transfer_sparse_codec.py`](../../nemo_rl/utils/weight_transfer_sparse_codec.py) | | Run the shared pipeline and S3 transport | [`weight_transfer_remote_sparse.py`](../../nemo_rl/utils/weight_transfer_remote_sparse.py) | | Run the ZeroMQ transport and relay | [`weight_transfer_zmq.py`](../../nemo_rl/utils/weight_transfer_zmq.py) | | Queue receiver work and expose endpoints | [`vllm_sparse_refit.py`](../../nemo_rl/models/generation/vllm/vllm_sparse_refit.py) | | Apply canonical updates through native loaders | [`vllm_sparse_delta.py`](../../nemo_rl/models/generation/vllm/vllm_sparse_delta.py) | -## Refit protocol +## Protocol -### Initialize the baseline +### Baseline and ownership The policy initialization task starts baseline construction as soon as its workers are ready, while the independent vLLM model load continues. `VllmRemoteSparseWeightSynchronizer.init_communicator()` discovers the receiver endpoints and then joins the prelaunched baseline before setup returns. The first rollout therefore does not enter a redundant weight sync or race an -unfinished snapshot with policy training. Conversion tasks are split into two -deterministic paths without changing Megatron Bridge: +unfinished snapshot with policy training. + +Conversion tasks are split into deterministic paths without changing Megatron +Bridge: -- Every conversion task keeps its source baseline in MCore layout. A stable name - hash assigns replicated dense tensors across their combined DP/CP and TP +- Every conversion task keeps its source baseline in MCore layout. A stable + name hash assigns replicated dense tensors across their combined DP/CP and TP replicas, and expert tensors across their expert-DP replicas. TP, PP, EP, and - ETP still contribute their unique shards. This keeps exactly one source copy - while sharing baseline scans and uploads across equivalent ranks. -- Exact direct, column, row, replicated, and gated mappings use that baseline - to produce HF-coordinate deltas without Bridge export. The decision is based - on the resolved Bridge mapping type, not a model name or weight suffix. + ETP still contribute their unique shards. This keeps one source copy while + sharing baseline scans and uploads across equivalent ranks. +- Exact direct, column, row, replicated, and gated mappings produce canonical + HF-coordinate deltas without Bridge export. The decision uses the resolved + Bridge mapping type, not a model name or parameter suffix. - Tasks with a transform also keep a canonical HF baseline, sharded across workers by a stable hash of the HF name. Stable ownership is required because - later refits export only changed residual tasks and therefore have a different - chunk sequence. + later refits export only changed residual tasks and therefore produce a + different chunk sequence. + +Changed flat locations on an affine mapping are projected with its shard +dimension, TP or ETP rank, and EP-global expert number. Attention output +projections, Mamba affine weights, norms, routers, shared experts, and other +exact mappings use this same model-family-agnostic path. + +Transformed, grouped, padded, tied, adapter, custom-postprocessed, custom FSDP, +and FP8-parameter tasks use the residual path. One integer flag per conversion +task is reduced across the policy world, and Bridge exports only globally +changed tasks when their dependencies are known. If one member of a grouped +export changes, the complete group is exported. A custom Bridge postprocessor +can have undeclared cross-task dependencies, so it retains the full residual +set whenever any residual task changes. Compound QKV, Mamba packing, +permutations, fused exports, and padded or tied embeddings therefore preserve +Bridge semantics. The local baselines contain one copy of each unique source element across the policy workers, rather than a full HF copy per DP replica. Residual tasks have -one additional canonical HF copy distributed across the workers. Baselines use +one additional canonical HF copy distributed across workers. Baselines use file-backed `torch.from_file` tensors by default; `NRL_REFIT_BASELINE_IN_MEMORY=1` keeps them in RAM. Local snapshotting and the residual Bridge export run concurrently. @@ -110,35 +127,23 @@ reserve one reusable GPU byte buffer large enough for the largest canonical tensor. This does not export HF values or mutate vLLM weights, and it removes scratch allocation from the first timed refit. -On a fresh run, vLLM already holds the shared checkpoint. Baseline construction -starts early and can overlap initial generation, so the redundant initial full -sync is skipped. On resume, both clusters must still start from the same HF -weight version; sparse refit does not reconstruct a rollout baseline from an -arbitrary training checkpoint. - -### Compare and encode deltas - -Every uniquely owned local tensor is copied to CPU and compared bytewise. For -an affine mapping, changed flat locations are mapped into the unsharded HF -tensor using the TP or ETP rank, shard dimension, and EP-global expert number. -This covers more than FFNs: attention output projections, Mamba affine weights, -norms, routers, shared experts, and other exact mappings use the same path. - -For non-affine tasks, one integer flag per conversion task is reduced across the -policy world. Bridge exports only globally changed tasks, after which the -canonical HF tracker computes the sparse payload. If one member of a grouped -export changes, the complete group is exported. A model-specific Bridge -postprocessor can have undeclared cross-task dependencies, so such bridges keep -the full residual task set whenever any residual task changes. Compound QKV, -Mamba packing, permutations, grouped or fused exports, padded or tied -embeddings, and other custom transformations therefore retain Bridge semantics. +On a fresh run, vLLM already holds the shared checkpoint, so the redundant +initial full sync is skipped. On resume, both clusters must still start from +the same HF weight version; sparse refit does not reconstruct a rollout +baseline from an arbitrary training checkpoint. + +### Compare and encode + +Every uniquely owned local tensor is copied to CPU and compared bytewise through +an integer view with the same element width. The payload contains only changed +locations and values. The comparison remains proportional to model size because +every unique source element is copied and scanned; Adam can therefore make it +CPU-bound even when the wire payload is sparse. Policy-local comparison removes full-tensor TP/EP gathers and PP broadcasts for -directly projectable weights, but comparison itself is still proportional to -model size: every unique local element is copied to CPU and scanned. Let `P` be -directly projectable bytes, `R` residual source bytes, `R_changed` the residual -HF tensors selected by the task flags, and `s` the element change fraction. The -leading work is approximately: +directly projectable weights. Let `P` be directly projectable bytes, `R` +residual source bytes, `R_changed` the residual HF tensors selected by task +flags, and `s` the element change fraction. The leading work is approximately: ```text old: Bridge(P + R) + HF D2H/scan(P + R) + wire(s(P + R)) @@ -146,94 +151,94 @@ new: local D2H/scan(P + R) + Bridge(R_changed) + HF scan(R_changed) + wire(s(P + R)) + one O(number_of_tasks) flag all-reduce ``` -When Adam changes at least one element in every residual tensor, -`R_changed` approaches `R`; the gain then comes from removing Bridge work for -`P`, not from task filtering. It helps less when local D2H or CPU scanning is -the bottleneck, most bytes use custom transformations, or `s` is high. The -reported changed percentage is computed from unique policy-local source -elements, so the extra detector does not inflate it with a second HF scan. - -For each assigned chunk, `DeltaCompressionTracker` finds changed flat -locations through an integer view with the same element width. It encodes -either the absolute new bits (`overwrite`) or new bits XOR baseline bits -(`xor`). Both encodings are dtype-blind and preserve FP8 bit patterns in +When Adam changes at least one element in every residual tensor, `R_changed` +approaches `R`. The gain then comes from removing Bridge work for `P`, not from +task filtering. It helps less when local D2H or CPU scanning is the bottleneck, +most bytes use custom transformations, or `s` is high. The reported changed +percentage is computed from unique policy-local source elements, so the residual +detector does not inflate it with a second HF scan. + +`DeltaCompressionTracker` finds changed flat locations through the equal-width +integer view. Both encodings are dtype-blind and preserve FP8 bit patterns in the codec, although end-to-end FP8 rollout refit is outside the supported scope. -`overwrite` is idempotent and recommended. `xor` can compress better, but it -requires an exact same-dtype receiver baseline and exactly-once application. -Selecting `xor` enables mixed operation: directly projected, bitwise-compatible -policy shards use XOR, while Bridge residuals and the full-HF compatibility -path use overwrite. A payload batch may therefore contain both operations. The -receiver validates direct-loader compatibility and fails closed on a transform, -dtype cast, or overlapping mapping; use `overwrite` for models with any such -direct loader. The pending source baseline always records the exact new source -bits. +`overwrite` carries absolute new bits, is idempotent, and is recommended. `xor` +carries new bits XOR baseline bits and may compress better, but requires exact +same-dtype state and exactly-once application. Selecting XOR creates +mixed-operation payloads: direct bitwise-compatible policy shards use XOR, +while Bridge residuals and the full-HF compatibility path use overwrite. A +payload batch can contain both operations. The receiver fails closed on a +transform, dtype cast, or overlapping XOR destination. Use overwrite when any +native loader cannot preserve those XOR requirements. The pending source +baseline always records the exact new source bits. The producer pulls bounded export chunks and compares them in parallel. A separate bounded stage coalesces encoded chunks up to `sparse_bucket_size_bytes`, serializes them, and applies zstd level 1 before the -transport executor. Separating the 256 MiB S3 compare chunk from the 1 GiB wire -bucket preserves D2H/scan parallelism while reducing object and manifest count. -The stages run concurrently, so payload N transfers while later chunks are -compared and encoded. Source baselines do not commit until the entire transfer -succeeds. Worker errors are reported only after every rank drains its Bridge -export iterator; stopping early can strand peers in a conversion collective. +transport executor. Keeping the 256 MiB compare chunk smaller than the +recommended 1 GiB wire bucket preserves D2H/scan parallelism while reducing +object and manifest count. Payload N can transfer while later chunks are +compared and encoded. Source baselines do not commit until the complete transfer +and receiver flush succeed. Pipeline errors cancel outstanding local futures +and propagate to the synchronizer. -### Transfer and apply +### Transport and apply -| Behavior | S3 | ZeroMQ | +| Property | S3 | ZeroMQ | |---|---|---| -| Value plane | AWS CRT `PUT_OBJECT` | DEALER to ROUTER relay | -| Receiver notification | HTTP object manifest | Relay HTTP fanout | -| Retry identity | Object key and checksum | Transfer, producer, payload IDs, and checksum | -| Lifetime | Delete after all receivers respond | No persistent object | +| Value plane | AWS CRT multipart `PUT_OBJECT` | DEALER to ROUTER relay | +| Notification | HTTP object manifest | Relay HTTP fanout | +| Retry identity | Object key + checksum | Transfer + producer + payload IDs + checksum | +| Lifetime | Delete after all receivers reply | No persistent object | -S3 uses 64 MiB multipart parts, a 2 GiB client memory limit, and a 10 Gbps CRT -throughput target. ZeroMQ assigns each producer to one relay; that relay fans -the compressed payload out to every generation replica. Both transports use -the same receiver endpoints and checksum validation. +S3 uses 64 MiB multipart parts, a 2 GiB CRT client memory limit, and a 10 Gbps +throughput target. ZeroMQ assigns each producer to one inference-cluster relay; +that relay fans compressed bytes to every generation replica and avoids +duplicate cross-cluster traffic. Both transports use the same receiver +endpoints and checksum validation. The receiver deduplicates payload identities and applies bounded batches on one -FIFO worker thread. Each generation replica downloads a transport payload -once. When its vLLM ranks share a node, decode and flat-file staging under -`/dev/shm` begin as soon as each payload arrives, without waiting for the batch -to fill. Staging futures then feed the serial collective apply worker, so -download, decompression, staging, and earlier GPU applies can overlap. Queue -depth provides backpressure; the default depth and batch size bound the pending -window at 256 payloads. +FIFO worker thread. Each generation replica downloads a transport payload once. +When its vLLM ranks share a node, decode and flat-file staging under `/dev/shm` +begin as soon as each payload arrives, without waiting for the batch to fill. +Staging futures feed the serial collective apply worker, so download, +decompression, staging, and earlier GPU applies can overlap. Queue depth limits +submitted work to 32 batches by default; with batches of eight, that is roughly +256 payloads plus the current partial batch. Locations use `int32` unless a single canonical tensor exceeds the signed -32-bit index range; values remain grouped by dtype. The collective RPC passes -only file paths, and workers use `torch.load(..., mmap=True)`. The mmap is not -a second baseline: it lets colocated ranks share the staged file's page cache -instead of materializing independent CPU copies. The staged format flattens -locations into one `int32` and one `int64` tensor; it does not serialize one -tensor object per model parameter. If ranks do not share a node, the receiver -still decodes once and sends that flat representation through one collective -RPC. - -There is deliberately no TP/EP source plan. Every worker sees the canonical -sparse entries, scatters them into its reusable dense source buffer, and lets -the native loader select the local destination. This removes persistent vLLM -placement knowledge from NeMo RL, at the cost of canonical-tensor GPU -initialization and duplicated sparse H2D across ranks. Measure that cost on the -target TP/EP topology; H2D no longer scales only with the worker-local sparse -subset. - -During the existing untimed metadata prewarm, one no-op native-loader pass -records names that issue no model-storage copy on that fixed rank. Later refits -skip scratch construction for those explicit pipeline/expert/MTP skips. The -cache contains names only; it stores no placement offsets, tensor routes, or +32-bit index range; values remain grouped by dtype. For shared-node workers the +collective RPC passes only staged file paths, and each rank uses +`torch.load(..., mmap=True)`. The mmap is not a second baseline: it lets ranks +share the staged file's page cache instead of materializing independent CPU +copies. The format flattens locations into one `int32` and one `int64` tensor; +it does not serialize one tensor object per model parameter. When ranks do not +share a node, the receiver still decodes once and sends the same flat +representation through one collective RPC. + +There is deliberately no TP/EP source plan. Every vLLM worker sees canonical +sparse entries, scatters them into its reusable dense source buffer, and calls +the model's native `load_weights()`. Native loaders own QKV, MoE, Mamba, TP, and +EP placement; NeMo RL stores no route model or family-specific placement +formulas and does not patch vLLM. This trades canonical-tensor scratch +initialization and duplicated sparse H2D across ranks for independence from vLLM +internals. Measure that cost on the target TP/EP topology; H2D no longer scales +only with the worker-local sparse subset. + +During the untimed metadata prewarm, one no-op native-loader pass records names +that issue no model-storage copy on that fixed rank. Later refits skip scratch +construction for those explicit pipeline, expert, or MTP skips. The cache +contains names only; it stores no placement offsets, tensor routes, or model-family rules. -The final `/nemo-rl/refit/flush` drains the queue, synchronizes CUDA, and checks -optional delta samples. Only then does the source commit pending baseline -updates in background CPU threads. +The final `/nemo-rl/refit/flush` drains every batch, synchronizes CUDA, and +checks optional delta samples. Only then does the source commit exact pending +baseline bits in background CPU threads. > **Failure boundary:** source baseline commit is transactional, but receiver -> updates are in place and are not rolled back. If a transfer fails after a +> writes are in place and are not rolled back. If a transfer fails after a > receiver accepts any payload, reload that receiver from a known-good weight -> version before retrying. This is mandatory for `xor`, because replaying an -> already-applied XOR reverts those bits. Replaying `overwrite` is safe. +> version before retrying. This is mandatory for XOR, because replaying an +> already-applied XOR reverts those bits. Replaying overwrite is safe. ## Payload and native apply @@ -245,35 +250,34 @@ Each serialized payload is: Contiguous locations use a range encoding. Other sorted locations are delta-encoded into the smallest lossless unsigned width among 16, 32, and 64 -bits. Metadata carries the HF name and shape, value offsets, location encoding, -the `xor` or `overwrite` operation, and an optional verification sample budget. +bits. Metadata carries the HF name and shape, dtype, value offsets, location +encoding, XOR or overwrite operation, and optional verification sample budget. HF coordinates are the canonical wire format because Megatron Bridge defines the training-to-HF mapping while vLLM owns the packed and sharded destination. -For each item, the receiver resets the resident largest-tensor scratch buffer, -scatters the sparse values, and calls the model's native `load_weights()`. +For each item, the receiver resets its resident largest-tensor scratch buffer, +scatters sparse values, and calls the model's native `load_weights()`. + A storage-scoped PyTorch dispatch mode changes only copies into model parameter -or buffer storage; it never encodes QKV, MoE, Mamba, TP, or EP geometry. - -For XOR, unchanged scratch bits are zero and target copies become bitwise XOR. -The source must remain a view of the scratch storage, dtypes must match, and -overlapping destination copies fail closed. For overwrite, unchanged entries -are NaN sentinels. The dispatch mode propagates the first sparse mask through -subsequent native copies and writes only selected destination entries. It keeps -only those entries for per-item rollback rather than cloning the full target. -This supports native pointwise transforms and dtype casts without model-specific -formulas. One-byte FP8 overwrite uses an exact bit sentinel and therefore -requires a non-transforming native loader; end-to-end quantized rollout refit -remains out of scope. +or buffer storage; it never encodes QKV, MoE, Mamba, TP, or EP geometry. For +XOR, unchanged scratch bits are zero and target copies become bitwise XOR. The +source must remain a view of scratch storage, dtypes must match, and overlapping +destination copies fail closed. For overwrite, unchanged entries are NaN +sentinels. The dispatch mode propagates the first sparse mask through subsequent +native copies and writes only selected destination entries. It keeps no target +backup because a partial batch already requires receiver reload. This supports +native pointwise transforms and dtype casts without model-specific formulas. +One-byte FP8 overwrite uses an exact bit sentinel and therefore requires a +non-transforming loader; end-to-end quantized rollout refit remains out of +scope. Native-loader return values distinguish an explicit skip from an unsupported -apply. An empty loaded set is accepted, matching vLLM's existing handling of -pipeline/expert ownership and checkpoint-only parameters such as inactive MTP -weights. Loader exceptions propagate. A loader that reports a weight loaded but -does not issue a supported target copy fails closed. There is no layout fallback -and no cached loader trace or route model. NeMo RL does not patch vLLM or encode -any vLLM layout: its only integration point is the model's public -`load_weights()` behavior. +apply. An empty loaded set is accepted, matching vLLM's handling of pipeline or +expert ownership and checkpoint-only parameters such as inactive MTP weights. +Loader exceptions propagate. A loader that reports a weight loaded without a +supported target copy fails closed. There is no layout fallback, cached loader +trace, or route model. NeMo RL's only integration point is the native model's +public `load_weights()` behavior. ## Configuration @@ -285,7 +289,7 @@ policy: backend: vllm refit_transport: vllm_s3_sparse # or vllm_zmq_sparse delta_compression: - encoding: overwrite # xor requires bitwise-compatible direct vLLM loaders + encoding: overwrite # xor requires bitwise-compatible native loaders sparse_bucket_size_bytes: 1073741824 colocated: enabled: false @@ -298,11 +302,11 @@ policy: zmq_refit_server_port: null ``` -S3 requires `NRL_REFIT_S3_BUCKET`; region and key prefix default to -`us-east-1` and `nemo-rl-refit`. ZeroMQ requires routable TCP access to the -relay port. The HTTP and ZeroMQ servers are plaintext, so use a trusted or -encrypted network. When `http_refit_api_key_env_var` is set, the named variable -must contain the same nonempty token on producers and receivers. +S3 requires `NRL_REFIT_S3_BUCKET`; region and key prefix default to `us-east-1` +and `nemo-rl-refit`. ZeroMQ requires routable TCP access to the relay port. The +HTTP and ZeroMQ servers are plaintext, so use a trusted or encrypted network. +When `http_refit_api_key_env_var` is set, the named variable must contain the +same nonempty token on producers and receivers. | Control | Default | |---|---:| @@ -319,10 +323,9 @@ must contain the same nonempty token on producers and receivers. | `NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD` | 0 | Export chunks are capped by `sparse_bucket_size_bytes` and the packed tensor -limit, but they intentionally remain smaller than the recommended S3 wire -bucket. Increase one concurrency control at a time; excessive parallelism can -move the bottleneck into host memory, collective export, relay fanout, or -receiver apply. +limit, but intentionally remain smaller than the recommended S3 wire bucket. +Increase one concurrency control at a time; excess parallelism moves the +bottleneck into host memory, Bridge export, relay fanout, or receiver apply. ## Metrics and profiling @@ -330,44 +333,38 @@ receiver apply. |---|---| | `REFIT_BASELINE_INIT` | Baseline export and snapshot time | | `REFIT_RECEIVER_PREWARM` | GPU scratch reservation and rank-local native skip discovery | -| `REFIT_{S3,ZMQ}_TIMING` | Producer wall time, stage service time, payloads, bytes, and changed density | +| `REFIT_{S3,ZMQ}_TIMING` | Producer wall time, stage service times, payloads, bytes, and density | | `REFIT_{S3,ZMQ}_DELTA_CHANGE` | Global changed and total element counts | -| `REFIT_RECEIVER_TIMING` | Receiver staging span/wait, batches, apply time, and verification counts | +| `REFIT_RECEIVER_TIMING` | Receiver staging span/wait, batches, apply, and verification | | `REFIT_{S3,ZMQ}_DELTA_VERIFY` | Sampled transmitted-delta accuracy | -| `REFIT_{S3,ZMQ}_GLOBAL_COMMIT` | Successful transfer flush | +| `REFIT_{S3,ZMQ}_GLOBAL_COMMIT` | Successful global flush | `total_s` is producer wall time. Stage fields such as `encode_s`, `s3_put_s`, and `zmq_send_s` are sums across concurrent tasks and can exceed `total_s`; do -not add them as serial phases. Receiver responses additionally expose node -decode/staging, worker deserialization, scratch preparation, and native-loader -apply time. -These are also concurrent sums; compare them with receiver wall time rather -than adding them. The `partition` field is `none` for uniquely owned -policy-local shards, `names` for stable name-sharded residual exports, and -`chunks` for the full-HF compatibility path. - -Benchmark and profiler reports are dated artifacts under `profiles/`. They must -record the exact commit, image, topology, model revision, changed density, -payload settings, per-stage overlap, and sampled correctness. Do not copy an -older report's fitted values into this document as current results. - -The synchronizer returns metrics under `refit/delta/*`, -`refit/delta_verify/*`, and `refit/transfer/*` when GRPO logs them. These are -available to W&B and other configured loggers. End-to-end refit latency is -reported as `timing/train/prepare_for_generation/transfer_and_update_weights`. - -For Nsight Systems, use the existing baseline, policy stream, and vLLM -sparse-apply NVTX ranges. Producer and receiver thread names begin with -`nrl-refit-`, `nrl-zmq-`, or `nrl-vllm-sparse-refit`. +not add them as serial phases. Receiver responses also expose node +decode/staging, worker deserialization, and native-loader apply time. These are +concurrent sums as well, so compare them with receiver wall time rather than +adding them. + +The `partition` field is `none` for uniquely owned policy-local shards, `names` +for stable name-sharded residual exports, and `chunks` for the full-HF +compatibility path. Synchronizer metrics appear under `refit/delta/*`, +`refit/delta_verify/*`, and `refit/transfer/*` in W&B and other configured +loggers. End-to-end latency is +`timing/train/prepare_for_generation/transfer_and_update_weights`. + +Nsight ranges cover baseline creation, policy streaming, and vLLM sparse apply. +Relevant thread names start with `nrl-refit-`, `nrl-zmq-`, or +`nrl-vllm-sparse-refit`. ## Development and validation Keep transport changes behind the shared `stream_sparse_delta_payloads()` -pipeline. A transport should provide payload delivery and timing only; it must -not duplicate the baseline tracker, codec, receiver queue, or apply logic. -Retries must preserve payload identity and bytes, fan out to every required -replica, and require a successful global flush before baseline commit. Never -retry XOR after an uncertain or partial receiver apply. +pipeline. A transport provides payload delivery and timing only; it must not +duplicate the baseline tracker, codec, receiver queue, or apply logic. Retries +must preserve payload identity and bytes, fan out to every required replica, +and require a successful global flush before source commit. Never retry XOR +after an uncertain or partial receiver apply. Do not add model-specific placement math or a persistent placement cache. New layouts must work through their native vLLM weight loader and the generic @@ -378,8 +375,8 @@ transformed XOR, and overlapping XOR copies must fail closed. Codec changes must update encoder and decoder together, preserve 64-bit-safe locations, and commit exact source bits only after global success. Receiver -changes must preserve FIFO application, bounded memory, deferred-error -propagation, flush, CUDA synchronization, and clean shutdown. +changes must preserve FIFO application, bounded memory, error propagation, +flush, CUDA synchronization, and clean shutdown. Every conversion task must remain represented by a unique local baseline unless the complete policy-local path is disabled for FP8 parameters, quantization, or @@ -389,9 +386,9 @@ tensor. Tests must cover column and row offsets, gated splits, replicated ownership, nonzero TP/ETP ranks, EP-global expert naming, DP/CP and expert-DP ownership, transactional baseline updates, global changed-task agreement, stable residual ownership under filtering, grouped-task expansion, and fallback -for transformed mappings. Do not infer an unknown mapping from its parameter -suffix, drop Bridge task dependencies, or modify Megatron Bridge to expose a -transport-specific hook. +for transformed mappings. Do not infer an unknown mapping from its suffix, drop +Bridge task dependencies, or modify Megatron Bridge for a transport-specific +hook. Run the focused suite: @@ -407,6 +404,7 @@ uv run --extra vllm pytest -q -m vllm \ uv run ruff check \ nemo_rl/utils/weight_transfer_{remote_sparse,sparse_codec,zmq}.py \ + nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py \ nemo_rl/models/generation/vllm/vllm_{sparse_refit,sparse_delta}.py \ nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py \ tools/refit_bandwidth_calculator.py @@ -416,21 +414,21 @@ On the target topology, verify the exact commit, image digest, and checkpoint revision; validate fresh starts and same-version resumes; compare two balanced repetitions with an equivalent NCCL or full control; and require the requested changed density, one global commit, no traceback, and zero sampled mismatches. -After failure injection, confirm the source baseline does not commit and reload -the receiver before retrying. +After failure injection, confirm that the source baseline does not commit and +reload the receiver before retrying. ## Refit bandwidth calculator [`refit_bandwidth_calculator.py`](../../tools/refit_bandwidth_calculator.py) is a -calibrated comparison of the checked-in S3 and ZeroMQ measurements against a -measured H100 NCCL envelope. It is not a general fabric or topology simulator. +calibrated comparison of the checked-in zstd S3 and ZeroMQ measurements against +a measured H100 NCCL envelope. It is not a general fabric or topology simulator. -The sparse side evaluates the latency fits in `_SPARSE_LATENCY_FITS` for the -requested model size, transport, compression, and any positive -`--changed-pct`. The 3% and 5% fits are measured calibration points; other -densities are explicit extrapolations. The coefficients model end-to-end -latency, not transport bandwidth, so `--candidate-ethernet-gbps` does not -rescale S3 or ZeroMQ. +The sparse side evaluates `_SPARSE_LATENCY_FITS` for the requested model size, +transport, and any positive `--changed-pct`. The 3% and 5% zstd fits are +measured calibration points; other densities are explicit extrapolations. The +coefficients model end-to-end latency, not transport bandwidth, so +`--candidate-ethernet-gbps` does not rescale S3 or ZeroMQ. Wire bytes use the +corresponding measured zstd size fit. The NCCL side interpolates `_NCCL_ANCHORS` in log model-size space. Those anchors were measured at 400 Gbps per rank and are projected onto the requested @@ -441,30 +439,27 @@ T_ethernet = T_H100_IB * 400 / candidate_ethernet_gbps ``` `--candidate-ethernet-gbps` is raw bandwidth per rank, not aggregate node or -cluster bandwidth. This makes NCCL and the candidate Ethernet refer to the same +cluster bandwidth. This makes NCCL and candidate Ethernet refer to the same per-rank link while leaving the independently measured sparse path unchanged. ```bash uv run python tools/refit_bandwidth_calculator.py \ --model-size-gb 247.2 \ --changed-pct 3 \ - --compression zstd \ --candidate-ethernet-gbps 25 ``` The output reports the reference and projected NCCL envelopes, sparse latency, -estimated wire bytes, and the per-rank Ethernet crossover. Below the lower +estimated wire bytes, and per-rank Ethernet crossover. Below the lower crossover sparse refit beats the complete NCCL envelope; above the upper crossover NCCL wins; between them the measured range has no single winner. `--json` emits the same fields for scripts. -The production transport currently applies zstd level 1 to every payload. -`--compression raw` selects the checked-in uncompressed calibration for -analysis; it is not a runtime switch. Production payloads use zstd level 1. -Treat values outside the calibration range, or a different topology and -parallel mapping, as experiment inputs rather than performance claims. Update -the constants only from a balanced profile matrix and keep the source artifact -under `profiles/`. +The production transport applies zstd level 1 to every payload, so the +calculator intentionally has no synthetic raw-compression arm. Treat values +outside the 63.2-1121 GB model range or 3%-5% density range, and any different +topology or parallel mapping, as experiment inputs rather than performance +claims. ## Failure guide @@ -472,8 +467,8 @@ under `profiles/`. |---|---| | Baseline is missing a tensor | Check baseline completion, checkpoint equality, and Bridge name mappings. | | No refit endpoint is found | Check worker startup, fixed ports, routing, and network policy. | -| Every worker reports a tensor unloaded | Verify the canonical HF name and native loader; do not add model-specific placement math. | +| Every worker reports a tensor unloaded | Verify the canonical HF name and native loader; do not add placement formulas. | | A payload ID is reused with different bytes | Start a new transfer or resend the original payload unchanged. | | Changed percentage rises unexpectedly | Correlate `DELTA_CHANGE` with `GLOBAL_COMMIT` and baseline commit completion. | -| Apply queue stalls | Inspect receiver timing and reduce source or relay concurrency. | +| Apply queue stalls | Inspect receiver timing; reduce producer or relay concurrency. | | A transfer fails after payload acceptance | Reload the receiver from a known-good checkpoint before retrying. | diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index 7292c40fa5d..4f7e0b58aa9 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -908,6 +908,8 @@ def _spinup_nemo_gym(base_urls, model_name): backend = generation_config["backend"] generation_config["model_name"] = policy_config["model_name"] # Needed for vLLM remote_transport = None + remote_synchronizer_cls = None + remote_baseline_init_refs: list[Any] = [] # Dictionary to store worker initialization timing stats for logging worker_init_timing_metrics = {} @@ -968,6 +970,11 @@ def init_policy(): init_optimizer=True, init_reference_model=init_reference_model, ) + if remote_transport is not None: + assert remote_synchronizer_cls is not None + remote_baseline_init_refs.extend( + remote_synchronizer_cls.start_baseline(p, remote_transport) + ) return p, time.perf_counter() - t0 def init_vllm(): @@ -1007,8 +1014,6 @@ def init_megatron_generation(policy=None): ) return mg, time.perf_counter() - t0 - init_policy_for_generation = init_policy - def initialize_generation_with_policy( init_generation_fn, generation_name: str, @@ -1042,7 +1047,7 @@ def initialize_generation_with_policy( parallel_start_time = time.perf_counter() with ThreadPoolExecutor(max_workers=2) as executor: generation_future = executor.submit(init_generation_fn) - policy_future = executor.submit(init_policy_for_generation) + policy_future = executor.submit(init_policy) policy_generation, generation_time = generation_future.result() policy, policy_time = policy_future.result() parallel_wall_time = time.perf_counter() - parallel_start_time @@ -1064,7 +1069,7 @@ def initialize_generation_with_policy( policy_generation, generation_time = init_generation_fn() worker_init_timing_metrics[init_time_key] = generation_time - policy, policy_time = init_policy_for_generation() + policy, policy_time = init_policy() worker_init_timing_metrics["policy_init_time_s"] = policy_time worker_init_timing_metrics["parallel_init_enabled"] = 0.0 @@ -1098,7 +1103,6 @@ def initialize_generation_with_policy( elif backend == "vllm": # vLLM generation: setup config, then initialize with policy generation_config = cast(VllmConfig, generation_config) - remote_baseline_init_refs: list[Any] = [] if generation_config.get("refit_transport") is not None: # Keep optional remote transport dependencies off the default path. from nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer import ( @@ -1112,15 +1116,7 @@ def initialize_generation_with_policy( megatron_enabled=policy_config["megatron_cfg"]["enabled"], ) assert remote_transport is not None - - def init_policy_for_generation(): - policy, policy_time = init_policy() - remote_baseline_init_refs.extend( - VllmRemoteSparseWeightSynchronizer.start_baseline( - policy, remote_transport - ) - ) - return policy, policy_time + remote_synchronizer_cls = VllmRemoteSparseWeightSynchronizer if generation_config["vllm_cfg"]["precision"] == "fp8": assert loss_config.use_importance_sampling_correction, ( @@ -1187,13 +1183,13 @@ def init_nemo_gym(): def init_vllm_then_policy(): pg, vllm_t = init_vllm_deferred() - p, policy_t = init_policy_for_generation() + p, policy_t = init_policy() return pg, vllm_t, p, policy_t init_tasks["vllm_policy"] = init_vllm_then_policy else: init_tasks["vllm"] = init_vllm_deferred - init_tasks["policy"] = init_policy_for_generation + init_tasks["policy"] = init_policy init_tasks["nemo_gym"] = init_nemo_gym print( @@ -1301,7 +1297,8 @@ def init_vllm_then_policy(): if remote_transport is not None: t0 = time.perf_counter() assert isinstance(policy_generation, VllmGeneration) - policy_generation.weight_synchronizer = VllmRemoteSparseWeightSynchronizer( + assert remote_synchronizer_cls is not None + policy_generation.weight_synchronizer = remote_synchronizer_cls( policy, policy_generation, transport=remote_transport, diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index 9a3e76d72ef..2dc2817c867 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -513,18 +513,10 @@ def update_weights_from_collective(self) -> bool: return True def update_weights_from_decoded_sparse_payload( - self, - *serialized_payloads: bytes, - ) -> dict[str, Any]: - applier = self._get_sparse_delta_applier() - return applier.update_weights_from_decoded_sparse_payload(*serialized_payloads) - - def update_weights_from_decoded_sparse_payload_files( - self, - *payload_paths: str, + self, *payloads: bytes | str ) -> dict[str, Any]: applier = self._get_sparse_delta_applier() - return applier.update_weights_from_decoded_sparse_payload_files(*payload_paths) + return applier.update_weights_from_decoded_sparse_payload(*payloads) def synchronize_device(self) -> None: self._get_sparse_delta_applier().synchronize_device() diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py index aa2c4277a81..e397c154365 100644 --- a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py +++ b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py @@ -26,19 +26,25 @@ from nemo_rl.utils import weight_transfer_sparse_codec as sparse_codec from nemo_rl.utils.nsys import wrap_with_nvtx_name +_TensorViewKey = tuple[int, int, tuple[int, ...], tuple[int, ...]] +_LoaderWeight = tuple[str, torch.Tensor, sparse_codec.SparseOperation, int, int | None] + def _storage_key(tensor: torch.Tensor) -> int: return tensor.untyped_storage()._cdata -def _integer_view(tensor: torch.Tensor) -> torch.Tensor: - return tensor.view( - sparse_codec.integer_dtype_for_element_size(tensor.element_size()) +def _view_key(tensor: torch.Tensor) -> _TensorViewKey: + return ( + _storage_key(tensor), + int(tensor.storage_offset()), + tuple(map(int, tensor.shape)), + tuple(map(int, tensor.stride())), ) class _SparseWeightLoadMode(TorchDispatchMode): - """Turn native loader copies into sparse XOR or transactional overwrite.""" + """Turn native loader copies into sparse XOR or overwrite.""" def __init__( self, @@ -52,13 +58,9 @@ def __init__( self._operation: sparse_codec.SparseOperation = "overwrite" self._sample_limit = 0 self._exact_sentinel: int | None = None - self._backups: list[tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = [] - self._active_masks: dict[ - tuple[int, int, tuple[int, ...], tuple[int, ...]], torch.Tensor - ] = {} + self._active_masks: dict[_TensorViewKey, torch.Tensor] = {} self._verification_masks: dict[ - tuple[int, int, tuple[int, ...], tuple[int, ...]], - tuple[torch.Tensor, torch.Tensor], + _TensorViewKey, tuple[torch.Tensor, torch.Tensor] ] = {} self._xor_spans: dict[int, list[tuple[int, int]]] = {} self.copies = 0 @@ -70,8 +72,6 @@ def start( sample_limit: int, exact_sentinel: int | None, ) -> None: - if self._backups: - raise RuntimeError("Previous sparse native-loader item was not finished.") self._source_storage = _storage_key(source) self._operation = operation self._sample_limit = sample_limit @@ -81,25 +81,12 @@ def start( self._xor_spans.clear() self.copies = 0 - def _overwrite_changed( - self, - destination: torch.Tensor, - changed: torch.Tensor, - values: torch.Tensor, - ) -> None: - backup = destination.masked_select(changed) - self._backups.append((destination, changed, backup)) - destination.masked_scatter_(changed, values.to(destination.dtype)) - def _remember_changed( self, destination: torch.Tensor, changed: torch.Tensor ) -> None: - view_key = ( - _storage_key(destination), - int(destination.storage_offset()), - tuple(int(size) for size in destination.shape), - tuple(int(stride) for stride in destination.stride()), - ) + if self._sample_limit <= 0: + return + view_key = _view_key(destination) previous = self._verification_masks.get(view_key) if previous is not None: changed = previous[1] | changed @@ -124,12 +111,7 @@ def __torch_dispatch__( if not source.dtype.is_floating_point: raise RuntimeError("Sparse overwrite requires a floating-point loader.") source = source.expand_as(destination) - view_key = ( - _storage_key(destination), - int(destination.storage_offset()), - tuple(int(size) for size in destination.shape), - tuple(int(stride) for stride in destination.stride()), - ) + view_key = _view_key(destination) if self._exact_sentinel is not None: if ( _storage_key(source) != self._source_storage @@ -138,13 +120,11 @@ def __torch_dispatch__( raise RuntimeError( "Exact FP8 overwrite cannot pass through a transforming loader." ) - source_bits = _integer_view(source) + source_bits = sparse_codec.integer_view(source) changed = source_bits.ne(self._exact_sentinel) - destination_bits = _integer_view(destination) - self._overwrite_changed( - destination_bits, - changed, - source_bits.masked_select(changed), + destination_bits = sparse_codec.integer_view(destination) + destination_bits.masked_scatter_( + changed, source_bits.masked_select(changed) ) self._remember_changed(destination, changed) return destination @@ -152,10 +132,8 @@ def __torch_dispatch__( if changed is None or _storage_key(source) == self._source_storage: changed = ~torch.isnan(source) self._active_masks[view_key] = changed - self._overwrite_changed( - destination, - changed, - source.masked_select(changed), + destination.masked_scatter_( + changed, source.masked_select(changed).to(destination.dtype) ) self._remember_changed(destination, changed) return destination @@ -182,41 +160,31 @@ def __torch_dispatch__( if any(span[0] <= other[1] and other[0] <= span[1] for other in spans): raise RuntimeError("XOR native loader produced overlapping target copies.") spans.append(span) - destination_bits = _integer_view(destination) - source_bits = _integer_view(source) + destination_bits = sparse_codec.integer_view(destination) + source_bits = sparse_codec.integer_view(source) changed = source_bits.ne(0) values = destination_bits.masked_select(changed).bitwise_xor( source_bits.masked_select(changed) ) - self._overwrite_changed(destination_bits, changed, values) + destination_bits.masked_scatter_(changed, values) self._remember_changed(destination, changed) return destination - def _record(self, target: torch.Tensor, changed: torch.Tensor) -> None: - if self._sample_limit <= 0 or not target.is_contiguous(): - return - locations = changed.reshape(-1).nonzero().reshape(-1)[: self._sample_limit] - if not locations.numel(): - return - target_bits = _integer_view(target).reshape(-1) - self._verification.append( - (target, locations, target_bits.index_select(0, locations).clone()) - ) - self._sample_limit -= locations.numel() - @torch.no_grad() def finish(self) -> None: - """Commit the sparse copies after the loader finishes all transforms.""" - for destination, changed in self._verification_masks.values(): - self._record(destination, changed) - self._backups.clear() - self._verification_masks.clear() - - @torch.no_grad() - def rollback(self) -> None: - for destination, changed, backup in reversed(self._backups): - destination.masked_scatter_(changed, backup) - self._backups.clear() + """Record bounded target samples after the loader finishes transforms.""" + for target, changed in self._verification_masks.values(): + if self._sample_limit <= 0: + break + if not target.is_contiguous(): + continue + locations = changed.reshape(-1).nonzero().reshape(-1)[: self._sample_limit] + if locations.numel(): + target_bits = sparse_codec.integer_view(target).reshape(-1) + self._verification.append( + (target, locations, target_bits.index_select(0, locations).clone()) + ) + self._sample_limit -= locations.numel() self._verification_masks.clear() @@ -231,10 +199,8 @@ def __init__(self, model_runner: Any, device: torch.device) -> None: _storage_key(tensor) for tensor in (*model.parameters(), *model.buffers()) } self._scratch = torch.empty(0, dtype=torch.uint8, device=device) - self._classified_names: set[str] = set() self._skipped_names: set[str] = set() self._verification: list[tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = [] - self._verification_candidates = 0 def prewarm( self, state_dict_info: Mapping[str, tuple[tuple[int, ...], torch.dtype]] @@ -256,61 +222,30 @@ def discover_native_skips( pending = [ (name, shape, dtype) for name, (shape, dtype) in state_dict_info.items() - if name not in self._classified_names and dtype.is_floating_point + if dtype.is_floating_point ] if not pending: return - mode = _SparseWeightLoadMode(self._target_storages, []) - observations: list[tuple[str, int]] = [] - - def weights() -> Iterator[tuple[str, torch.Tensor]]: - active_name = None + def weights() -> Iterator[_LoaderWeight]: for name, shape, dtype in pending: - if active_name is not None: - mode.finish() - observations.append((active_name, mode.copies)) source = self._source_tensor( - { - "name": name, - "shape": shape, - "dtype": str(dtype).removeprefix("torch."), - } + {"shape": shape, "dtype": str(dtype).removeprefix("torch.")} ) source.fill_(float("nan")) exact_sentinel = ( - int(_integer_view(source).reshape(-1)[0].item()) + int(sparse_codec.integer_view(source).reshape(-1)[0].item()) if source.element_size() == 1 else None ) - mode.start(source, "overwrite", 0, exact_sentinel) - active_name = name - yield name, source - if active_name is not None: - mode.finish() - observations.append((active_name, mode.copies)) - - loader_weights = weights() - try: - with torch.no_grad(), mode: - loaded = self.model_runner.model.load_weights(loader_weights) - except Exception: - mode.rollback() - raise - finally: - loader_weights.close() + yield name, source, "overwrite", 0, exact_sentinel + loaded, observations = self._load_weights(weights(), []) if len(observations) != len(pending): raise RuntimeError( "Native loader did not consume all sparse weight metadata." ) - copied = sum(copies > 0 for _, copies in observations) - if loaded is not None and len(loaded) > copied: - raise RuntimeError( - "Native loader reported a loaded sparse weight without a " - "supported target copy." - ) - self._classified_names.update(name for name, _ in observations) + self._validate_loader_report(loaded, observations, allow_unknown_skips=True) if loaded is not None: self._skipped_names.update( name for name, copies in observations if copies == 0 @@ -321,7 +256,9 @@ def _source_tensor(self, item: dict[str, Any]) -> torch.Tensor: dtype = sparse_codec.dtype_from_name(str(item["dtype"])) byte_count = prod(shape) * dtype.itemsize if byte_count > self._scratch.numel(): - self.prewarm({str(item["name"]): (shape, dtype)}) + self._scratch = torch.empty( + byte_count, dtype=torch.uint8, device=self._scratch.device + ) return self._scratch[:byte_count].view(dtype).view(shape) @staticmethod @@ -331,7 +268,7 @@ def _scatter_values( locations: torch.Tensor, values: torch.Tensor, ) -> None: - source_bits = _integer_view(source).reshape(-1) + source_bits = sparse_codec.integer_view(source).reshape(-1) expected_dtype = sparse_codec.integer_dtype_for_element_size( source.element_size() ) @@ -356,8 +293,7 @@ def _prepare_loader_weight( item: dict[str, Any], locations: torch.Tensor, values: torch.Tensor, - mode: _SparseWeightLoadMode, - ) -> tuple[str, torch.Tensor]: + ) -> _LoaderWeight: operation = sparse_codec.sparse_operation(item["operation"]) source = self._source_tensor(item) exact_sentinel = None @@ -368,7 +304,7 @@ def _prepare_loader_weight( else: source.fill_(float("nan")) if source.element_size() == 1: - source_bits = _integer_view(source) + source_bits = sparse_codec.integer_view(source) exact_sentinel = int(source_bits.reshape(-1)[0].item()) if bool(values.eq(exact_sentinel).any()): exact_sentinel ^= 0x80 @@ -379,91 +315,99 @@ def _prepare_loader_weight( source_bits.fill_(exact_sentinel) self._scatter_values(source, item, locations, values) - sample_limit = int(item.get("verification_samples", 0)) - self._verification_candidates += sample_limit - mode.start( + return ( + str(item["name"]), source, operation, - sample_limit, + int(item.get("verification_samples", 0)), exact_sentinel, ) - return str(item["name"]), source - def _apply_decoded_items( + def _load_weights( self, - items: Iterable[tuple[dict[str, Any], torch.Tensor, torch.Tensor]], - ) -> None: - mode = _SparseWeightLoadMode(self._target_storages, self._verification) + weights: Iterable[_LoaderWeight], + verification: list[tuple[torch.Tensor, torch.Tensor, torch.Tensor]], + ) -> tuple[Any, list[tuple[str, int]]]: + mode = _SparseWeightLoadMode(self._target_storages, verification) yielded_names: list[str] = [] - copy_counts: list[int] = [] + observations: list[tuple[str, int]] = [] - def weights() -> Iterator[tuple[str, torch.Tensor]]: + def observed_weights() -> Iterator[tuple[str, torch.Tensor]]: active = False - for item, locations, values in items: - if str(item["name"]) in self._skipped_names: - continue + for name, source, operation, sample_limit, exact_sentinel in weights: if active: mode.finish() - copy_counts.append(mode.copies) - weight = self._prepare_loader_weight(item, locations, values, mode) + observations.append((yielded_names[-1], mode.copies)) + mode.start(source, operation, sample_limit, exact_sentinel) active = True - yielded_names.append(weight[0]) - yield weight + yielded_names.append(name) + yield name, source if active: mode.finish() - copy_counts.append(mode.copies) - - loader_weights = weights() - try: - with torch.no_grad(), mode: - loaded = self.model_runner.model.load_weights(loader_weights) - if len(copy_counts) != len(yielded_names): - raise RuntimeError("Native loader did not consume all sparse weights.") - copied_items = sum(copies > 0 for copies in copy_counts) - if loaded is None and copied_items != len(yielded_names): + observations.append((yielded_names[-1], mode.copies)) + + with torch.no_grad(), mode: + loaded = self.model_runner.model.load_weights(observed_weights()) + if len(observations) != len(yielded_names): + raise RuntimeError("Native loader did not consume all sparse weights.") + return loaded, observations + + @staticmethod + def _validate_loader_report( + loaded: Any, + observations: list[tuple[str, int]], + *, + allow_unknown_skips: bool, + ) -> None: + copied = sum(copies > 0 for _, copies in observations) + if loaded is None: + if not allow_unknown_skips and copied != len(observations): raise RuntimeError( "Native loader did not report whether uncopied sparse weights " "were skipped." ) - if loaded is not None and len(loaded) > copied_items: - raise RuntimeError( - "Native loader reported a loaded sparse weight without a " - "supported target copy." - ) - except Exception: - mode.rollback() - raise - finally: - loader_weights.close() + elif len(loaded) > copied: + raise RuntimeError( + "Native loader reported a loaded sparse weight without a supported " + "target copy." + ) - def _apply_decoded_item( + def _apply_decoded_items( self, - item: dict[str, Any], - locations: torch.Tensor, - values: torch.Tensor, + items: Iterable[tuple[dict[str, Any], torch.Tensor, torch.Tensor]], ) -> None: - self._apply_decoded_items(((item, locations, values),)) + def weights() -> Iterator[_LoaderWeight]: + for item, locations, values in items: + if str(item["name"]) in self._skipped_names: + continue + yield self._prepare_loader_weight(item, locations, values) + + loaded, observations = self._load_weights(weights(), self._verification) + self._validate_loader_report(loaded, observations, allow_unknown_skips=False) @wrap_with_nvtx_name( "vllm_internal_worker_extension/update_weights_from_decoded_sparse_payload" ) def update_weights_from_decoded_sparse_payload( - self, *serialized_payloads: bytes + self, *payloads: bytes | str ) -> dict[str, Any]: return self._load_decoded_sparse_payloads( - tuple(io.BytesIO(payload) for payload in serialized_payloads) + tuple( + io.BytesIO(payload) if isinstance(payload, bytes) else payload + for payload in payloads + ) ) def _load_decoded_sparse_payloads( self, sources: tuple[str | io.BytesIO, ...] ) -> dict[str, Any]: started = time.perf_counter() - deserialize_s = 0.0 - payloads: list[sparse_codec.DecodedSparsePayload] = [] - for source in sources: - item_started = time.perf_counter() - payloads.append( - cast( + deserialize_s = [0.0] + + def decoded_items() -> Iterator[sparse_codec.DecodedSparseItem]: + for source in sources: + item_started = time.perf_counter() + payload = cast( sparse_codec.DecodedSparsePayload, torch.load( source, @@ -472,41 +416,19 @@ def _load_decoded_sparse_payloads( mmap=isinstance(source, str), ), ) - ) - deserialize_s += time.perf_counter() - item_started + deserialize_s[0] += time.perf_counter() - item_started + yield from sparse_codec.iter_decoded_sparse_payload(payload) item_started = time.perf_counter() - self.prewarm( - { - str(item["name"]): ( - tuple(item["shape"]), - sparse_codec.dtype_from_name(str(item["dtype"])), - ) - for _, _, items in payloads - for item in items - } - ) - scratch_s = time.perf_counter() - item_started - item_started = time.perf_counter() - self._apply_decoded_items( - decoded_item - for payload in payloads - for decoded_item in sparse_codec.iter_decoded_sparse_payload(payload) - ) + self._apply_decoded_items(decoded_items()) sparse_apply_s = time.perf_counter() - item_started return { "ok": True, - "receiver_deserialize_s": deserialize_s, - "receiver_scratch_s": scratch_s, + "receiver_deserialize_s": deserialize_s[0], "receiver_sparse_apply_s": sparse_apply_s, "receiver_total_s": time.perf_counter() - started, } - def update_weights_from_decoded_sparse_payload_files( - self, *payload_paths: str - ) -> dict[str, Any]: - return self._load_decoded_sparse_payloads(payload_paths) - def synchronize_device(self) -> None: if torch.cuda.is_available(): torch.cuda.synchronize(self._cuda_device_index) @@ -515,7 +437,6 @@ def finish_sparse_delta_refit(self) -> dict[str, Any]: """Synchronize and compare bounded samples of target entries just changed.""" self.synchronize_device() verification, self._verification = self._verification, [] - candidates, self._verification_candidates = self._verification_candidates, 0 stats = ( torch.zeros(4, device=verification[0][0].device) if verification else None ) @@ -523,7 +444,9 @@ def finish_sparse_delta_refit(self) -> dict[str, Any]: with torch.no_grad(): for target, locations, expected_bits in verification: actual_bits = ( - _integer_view(target).reshape(-1).index_select(0, locations) + sparse_codec.integer_view(target) + .reshape(-1) + .index_select(0, locations) ) bit_mismatches = actual_bits.ne(expected_bits) actual = actual_bits.view(target.dtype).float() @@ -548,7 +471,7 @@ def finish_sparse_delta_refit(self) -> dict[str, Any]: values = [0.0] * 4 if stats is None else stats.cpu().tolist() return { "ok": True, - "verification_candidates": candidates, + "verification_candidates": samples, "verification_samples": samples, "verification_exact_mismatches": int(values[2]), "verification_mismatches": int(values[3]), diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py index 0d52c21fb05..faacb4d1d1e 100644 --- a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py +++ b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py @@ -202,7 +202,7 @@ def _enqueue_sparse_payload_apply( self._submit_pending_sparse_payloads() return response - def _submit_pending_sparse_payloads(self) -> Future[dict[str, Any]]: + def _submit_pending_sparse_payloads(self) -> None: payloads = tuple(self._refit_apply_pending_payloads) self._refit_apply_pending_payloads.clear() apply = ( @@ -213,7 +213,6 @@ def _submit_pending_sparse_payloads(self) -> Future[dict[str, Any]]: future = self._refit_apply_executor.submit(cast(Any, apply), payloads) self._refit_apply_futures.append(future) future.add_done_callback(self._notify_refit_apply_waiters) - return future def _notify_refit_apply_waiters(self, _future: Future[dict[str, Any]]) -> None: with self._refit_apply_queue_condition: @@ -306,7 +305,7 @@ def update_weights_from_staged_sparse_payloads( try: response = self._refit_collective_response( self._refit_collective_rpc( - "update_weights_from_decoded_sparse_payload_files", + "update_weights_from_decoded_sparse_payload", tuple(payload.path for payload in staged), ) ) diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 48bef53d731..a08bc0c6596 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -1840,7 +1840,9 @@ def _require_remote_sparse_refit(self) -> Any: MegatronRemoteSparseRefit, ) - self._remote_sparse_refit = MegatronRemoteSparseRefit.from_worker(self) + self._remote_sparse_refit = MegatronRemoteSparseRefit( + self, self.cfg["generation"]["delta_compression"] + ) return self._remote_sparse_refit def finish_remote_sparse_delta_sync(self, *, succeeded: bool) -> None: diff --git a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py index 398836cfe0e..762debd8538 100644 --- a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py +++ b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py @@ -23,7 +23,6 @@ import torch from nemo_rl.utils.weight_transfer_remote_sparse import ( - SparseDeltaStreamResult, init_sparse_delta_baseline_from_iterator, sparse_name_shard, stream_sparse_delta_payloads_via_s3_manifest, @@ -38,25 +37,18 @@ _COLUMN = 1 _ROW = 2 _REPLICATED = 3 -_DIRECT = 4 -_GATED = 5 +_GATED = 4 class MegatronRemoteSparseRefit: - @classmethod - def from_worker(cls, worker: Any) -> "MegatronRemoteSparseRefit": - generation_config = worker.cfg.get("generation") or {} - delta_config = generation_config.get("delta_compression") - if generation_config.get("refit_transport") is None or not delta_config: - raise RuntimeError("Remote sparse refit is not enabled for this worker.") - return cls(worker, delta_config) - def __init__(self, worker: Any, delta_config: Mapping[str, Any]) -> None: self._worker = worker self._delta_config = delta_config residual_config = dict(delta_config) if residual_config["encoding"] == "xor": residual_config["encoding"] = "overwrite" + # Residual Bridge outputs need a canonical-HF overwrite baseline; the + # optional policy tracker owns MCore-local projections and can retain XOR. self._tracker = DeltaCompressionTracker(residual_config) self._policy_tracker: DeltaCompressionTracker | None = None self._local_tensors: list[tuple[str, torch.Tensor]] = [] @@ -79,7 +71,7 @@ def _bridge_mapping_types() -> tuple[Any, dict[Any, int]]: return AutoMapping, { ColumnParallelMapping: _COLUMN, - DirectMapping: _DIRECT, + DirectMapping: _REPLICATED, GatedMLPMapping: _GATED, ReplicatedMapping: _REPLICATED, RowParallelMapping: _ROW, @@ -172,7 +164,7 @@ def _owns_policy_local_task(task: Any, *, replicated: bool = False) -> bool: @classmethod def _task_ownership(cls, task: Any, kind: int) -> tuple[torch.Tensor | None, bool]: tensor = task.param_weight - replicated = kind in (_DIRECT, _REPLICATED) or ( + replicated = kind == _REPLICATED or ( kind == _ROW and tensor is not None and tensor.ndim == 1 ) return ( @@ -228,10 +220,13 @@ def _projection( if tensor.ndim <= shard_dim: raise ValueError(f"Cannot shard {name!r} on dimension {shard_dim}.") global_shape = list(tensor.shape) - offsets = [0] * tensor.ndim global_shape[shard_dim] *= shard_count - offsets[shard_dim] = tensor.shape[shard_dim] * shard_rank - return SparseShardProjection(name, tuple(global_shape), tuple(offsets)) + return SparseShardProjection( + name, + tuple(global_shape), + shard_dim, + tensor.shape[shard_dim] * shard_rank, + ) @classmethod def _task_local_tensors( @@ -279,7 +274,7 @@ def _task_local_tensors( name = cls._canonical_hf_name(task, cast(str, hf_param)) projection = ( - SparseShardProjection(name, tuple(tensor.shape), (0,) * tensor.ndim) + SparseShardProjection(name, tuple(tensor.shape)) if replicated else cls._projection( name, @@ -440,7 +435,7 @@ def stream( timeout_s: float, shard_rank: int, shard_count: int, - ) -> SparseDeltaStreamResult: + ) -> dict[str, int]: self._prepare_paths() streamer = { "s3": stream_sparse_delta_payloads_via_s3_manifest, @@ -486,7 +481,7 @@ def stream( if local_future is not None else {"payloads": 0, "changed_elements": 0, "total_elements": 0} ) - result = SparseDeltaStreamResult( + result = dict( payloads=int(local_result["payloads"]) + int(misc_result["payloads"]), changed_elements=int(local_result["changed_elements"]) + misc_changed, total_elements=int(local_result["total_elements"]) + misc_total, diff --git a/nemo_rl/utils/weight_transfer_remote_sparse.py b/nemo_rl/utils/weight_transfer_remote_sparse.py index 8c084a84300..8b6de5930fc 100644 --- a/nemo_rl/utils/weight_transfer_remote_sparse.py +++ b/nemo_rl/utils/weight_transfer_remote_sparse.py @@ -24,7 +24,7 @@ from contextlib import suppress from dataclasses import dataclass from functools import cache -from typing import Any, Literal, TypedDict +from typing import Any, Literal from urllib.parse import quote import requests @@ -50,12 +50,6 @@ _S3_MEMORY_LIMIT = 2 * 1024**3 -class SparseDeltaStreamResult(TypedDict): - payloads: int - changed_elements: int - total_elements: int - - SparsePartitionMode = Literal["chunks", "names", "none"] @@ -230,9 +224,11 @@ def vllm_refit_endpoints(base_urls: Sequence[str], path: str) -> list[str]: def _partition_sparse_weights_by_name( tensors: Iterable[NamedTensor], shard_rank: int, shard_count: int ) -> Iterator[NamedTensor]: - for name, tensor in tensors: - if sparse_name_shard(name, shard_count) == shard_rank: - yield name, tensor + return ( + (name, tensor) + for name, tensor in tensors + if sparse_name_shard(name, shard_count) == shard_rank + ) def refit_http_session() -> requests.Session: @@ -336,7 +332,7 @@ def stream_sparse_delta_payloads( shard_rank: int, shard_count: int, partition: SparsePartitionMode = "chunks", -) -> SparseDeltaStreamResult: +) -> dict[str, int]: if partition == "names": iterator = _partition_sparse_weights_by_name(iterator, shard_rank, shard_count) prefix = transport.upper() @@ -408,29 +404,28 @@ def transfer_payload( } chunk_count = 0 export_pull_s = 0.0 - encode_inflight: set[Any] = set() + encode_inflight: dict[Any, None] = {} serialize_inflight: dict[Any, int] = {} - transfer_inflight: set[Any] = set() - worker_errors: list[Exception] = [] + transfer_inflight: dict[Any, None] = {} max_encode_inflight = encode_workers * 2 max_serialize_inflight = serialize_workers * 2 max_transfer_inflight = transfer_workers * 2 bucket = _SparsePayloadBucket([]) - def collect_transfers(*, block: bool) -> None: - if not transfer_inflight: + def resolved(inflight: dict[Any, Any], *, block: bool) -> Iterator[tuple[Any, Any]]: + if not inflight: return - if block: - completed, _ = wait(transfer_inflight, return_when=FIRST_COMPLETED) - else: - completed = {future for future in transfer_inflight if future.done()} + completed = ( + wait(inflight, return_when=FIRST_COMPLETED)[0] + if block + else tuple(future for future in inflight if future.done()) + ) for future in completed: - transfer_inflight.remove(future) - try: - result = future.result() - except Exception as error: - worker_errors.append(error) - continue + metadata = inflight.pop(future) + yield metadata, future.result() + + def collect_transfers(*, block: bool) -> None: + for _, result in resolved(transfer_inflight, block=block): counts["payloads"] += 1 counts["wire_bytes"] += int(result["body_size"]) for key, value in result.items(): @@ -441,24 +436,12 @@ def collect_transfers(*, block: bool) -> None: ) def collect_serialized(*, block: bool) -> None: - if not serialize_inflight: - return - if block: - completed, _ = wait(serialize_inflight, return_when=FIRST_COMPLETED) - else: - completed = {future for future in serialize_inflight if future.done()} - for future in completed: - index = serialize_inflight.pop(future) - try: - encoded = future.result() - except Exception as error: - worker_errors.append(error) - continue + for index, encoded in resolved(serialize_inflight, block=block): while len(transfer_inflight) >= max_transfer_inflight: collect_transfers(block=True) - transfer_inflight.add( + transfer_inflight[ transfer_executor.submit(transfer_payload, encoded, index) - ) + ] = None collect_transfers(block=False) def submit_bucket() -> None: @@ -496,13 +479,8 @@ def consume_encoded(encoded: Any) -> None: submit_bucket() def drain_encodes() -> None: - completed, _ = wait(encode_inflight, return_when=FIRST_COMPLETED) - for future in completed: - encode_inflight.remove(future) - try: - consume_encoded(future.result()) - except Exception as error: - worker_errors.append(error) + for _, encoded in resolved(encode_inflight, block=True): + consume_encoded(encoded) stream_start = time.perf_counter() try: @@ -515,7 +493,7 @@ def drain_encodes() -> None: continue while len(encode_inflight) >= max_encode_inflight: drain_encodes() - encode_inflight.add(encode_executor.submit(encode_chunk, chunk)) + encode_inflight[encode_executor.submit(encode_chunk, chunk)] = None while encode_inflight: drain_encodes() @@ -524,8 +502,6 @@ def drain_encodes() -> None: collect_serialized(block=True) while transfer_inflight: collect_transfers(block=True) - if worker_errors: - raise worker_errors[0] except Exception: for futures in (encode_inflight, serialize_inflight, transfer_inflight): for future in futures: @@ -576,7 +552,7 @@ def stream_sparse_delta_payloads_via_s3_manifest( shard_rank: int, shard_count: int, partition: SparsePartitionMode = "chunks", -) -> SparseDeltaStreamResult: +) -> dict[str, int]: urls = [url.strip().rstrip("/") for url in refit_targets if url.strip()] if not urls: raise ValueError("At least one vLLM S3 refit URL is required.") @@ -668,51 +644,16 @@ def post(url: str) -> dict[str, Any]: return result pool = executor or _executor("refit-fanout", len(endpoint_urls)) - futures = [pool.submit(post, url) for url in endpoint_urls] - return [future.result() for future in futures] - - -def flush_vllm_refit_urls( - base_urls: Sequence[str], - *, - api_key_env_var: str | None, - timeout_s: float, -) -> list[dict[str, Any]]: - return post_vllm_refit_endpoints( - vllm_refit_endpoints(base_urls, G_VLLM_REFIT_FLUSH_PATH), - {}, - api_key=vllm_refit_api_key(api_key_env_var), - timeout_s=timeout_s, - ) - - -def prepare_vllm_sparse_refit_urls( - base_urls: Sequence[str], - state_dict_info: Mapping[str, tuple[tuple[int, ...], torch.dtype]], - *, - api_key_env_var: str | None, - timeout_s: float, -) -> list[dict[str, Any]]: - tensors = { - name: [list(shape), str(dtype).removeprefix("torch.")] - for name, (shape, dtype) in state_dict_info.items() - } - return post_vllm_refit_endpoints( - vllm_refit_endpoints(base_urls, G_VLLM_REFIT_PREPARE_PATH), - {"tensors": tensors}, - api_key=vllm_refit_api_key(api_key_env_var), - timeout_s=timeout_s, - ) + return list(pool.map(post, endpoint_urls)) def download_s3_refit_payload( manifest: Mapping[str, Any], ) -> bytes: - bucket, region, key, checksum = ( - str(manifest[field]) for field in ("bucket", "region", "key", "checksum") + body = _get_manifest_s3_store(str(manifest["bucket"]), str(manifest["region"])).get( + str(manifest["key"]) ) - body = _get_manifest_s3_store(bucket, region).get(key) - return decode_sparse_payload(body, checksum) + return decode_sparse_payload(body, str(manifest["checksum"])) def merge_vllm_refit_metrics( diff --git a/nemo_rl/utils/weight_transfer_sparse_codec.py b/nemo_rl/utils/weight_transfer_sparse_codec.py index 6199a78eca1..72caa89e3f8 100644 --- a/nemo_rl/utils/weight_transfer_sparse_codec.py +++ b/nemo_rl/utils/weight_transfer_sparse_codec.py @@ -92,36 +92,24 @@ class SparseShardProjection: name: str global_shape: tuple[int, ...] - offsets: tuple[int, ...] + shard_dim: int | None = None + offset: int = 0 def map_locations( self, locations: torch.Tensor, local_shape: tuple[int, ...] ) -> torch.Tensor: - if len(local_shape) != len(self.global_shape) or len(local_shape) != len( - self.offsets - ): - raise ValueError(f"Sparse shard {self.name!r} has inconsistent ranks.") - if any( - offset < 0 or offset + local > global_size - for local, global_size, offset in zip( - local_shape, self.global_shape, self.offsets, strict=True - ) - ): - raise ValueError(f"Sparse shard {self.name!r} exceeds its global shape.") - if local_shape == self.global_shape and not any(self.offsets): + if self.shard_dim is None: return locations - - mapped = torch.zeros_like(locations) - for dim, (local_size, global_size, offset) in enumerate( - zip(local_shape, self.global_shape, self.offsets, strict=True) - ): - local_stride = prod(local_shape[dim + 1 :]) - global_stride = prod(self.global_shape[dim + 1 :]) - coordinate = torch.div( - locations, local_stride, rounding_mode="floor" - ).remainder(local_size) - mapped.add_((coordinate + offset) * global_stride) - return mapped + if self.shard_dim == 0: + return locations + self.offset * prod(local_shape[1:]) + inner = prod(local_shape[self.shard_dim + 1 :]) + local_slab = local_shape[self.shard_dim] * inner + outer = torch.div(locations, local_slab, rounding_mode="floor") + return ( + locations + + self.offset * inner + + outer * (prod(self.global_shape[self.shard_dim :]) - local_slab) + ) def integer_dtype_for_element_size(element_size: int) -> torch.dtype: @@ -145,10 +133,8 @@ def sparse_operation(value: object) -> SparseOperation: raise ValueError(f"Unsupported sparse-refit operation {value!r}.") -def _integer_view(tensor: torch.Tensor) -> torch.Tensor: - return tensor.contiguous().view( - integer_dtype_for_element_size(tensor.element_size()) - ) +def integer_view(tensor: torch.Tensor) -> torch.Tensor: + return tensor.view(integer_dtype_for_element_size(tensor.element_size())) def _bytewise_diff_mask(current: torch.Tensor, baseline: torch.Tensor) -> torch.Tensor: @@ -156,7 +142,7 @@ def _bytewise_diff_mask(current: torch.Tensor, baseline: torch.Tensor) -> torch. raise ValueError( "Current tensor and baseline must have identical shape and dtype." ) - return _integer_view(current) != _integer_view(baseline) + return integer_view(current) != integer_view(baseline) def encode_sparse_infos( @@ -351,15 +337,12 @@ def __init__( def prepare_sparse_delta_payload( self, tensors: TensorBatch ) -> PreparedTensorPayload: - self._wait_for_baseline_commits() - sparse_infos = [] - pending_updates = {} + sparse_infos: list[SparseInfo] = [] changed_elements = total_elements = 0 - for name, tensor in tensors: - baseline, current, locations, current_values = self._find_changes( - name, tensor - ) - baseline_bits = _integer_view(baseline).view(-1) + for name, baseline, current, locations, current_values in self._changes( + tensors + ): + baseline_bits = integer_view(baseline).view(-1) total_elements += current.numel() changed_elements += locations.numel() if locations.numel(): @@ -390,9 +373,6 @@ def prepare_sparse_delta_payload( self.encoding, ) ) - pending_updates[name] = (locations, current_values) - with self._pending_updates_lock: - self._pending_updates.update(pending_updates) payload = encode_sparse_infos(sparse_infos) if self.verification_samples: self._add_verification_samples(payload[2]) @@ -402,31 +382,34 @@ def prepare_change_summary( self, tensors: Iterable[NamedTensor] ) -> tuple[set[str], int, int]: """Scan local tensors without constructing a wire payload.""" - self._wait_for_baseline_commits() changed_names = set() - pending_updates = {} changed_elements = total_elements = 0 - for name, tensor in tensors: - _, current, locations, current_values = self._find_changes(name, tensor) + for name, _, current, locations, _ in self._changes(tensors): total_elements += current.numel() changed_elements += locations.numel() if locations.numel(): changed_names.add(name) - pending_updates[name] = (locations, current_values) - with self._pending_updates_lock: - self._pending_updates.update(pending_updates) return changed_names, changed_elements, total_elements - def _find_changes( - self, name: str, tensor: torch.Tensor - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - baseline = self.baseline.get(name) - if baseline is None: - raise RuntimeError(f"Sparse delta baseline is missing {name!r}.") - current = tensor.detach().cpu().contiguous() - locations = _bytewise_diff_mask(current, baseline).view(-1).nonzero().view(-1) - current_values = _integer_view(current).view(-1)[locations] - return baseline, current, locations, current_values + def _changes( + self, tensors: Iterable[NamedTensor] + ) -> Iterable[tuple[str, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]]: + self._wait_for_baseline_commits() + pending_updates = {} + for name, tensor in tensors: + baseline = self.baseline.get(name) + if baseline is None: + raise RuntimeError(f"Sparse delta baseline is missing {name!r}.") + current = tensor.detach().cpu().contiguous() + locations = ( + _bytewise_diff_mask(current, baseline).view(-1).nonzero().view(-1) + ) + values = integer_view(current).view(-1)[locations] + if locations.numel(): + pending_updates[name] = (locations, values) + yield name, baseline, current, locations, values + with self._pending_updates_lock: + self._pending_updates.update(pending_updates) def _add_verification_samples( self, @@ -483,7 +466,7 @@ def _commit_baseline_updates( updates: Iterable[tuple[str, tuple[torch.Tensor, torch.Tensor]]], ) -> None: for name, (locations, values) in updates: - target = _integer_view(self.baseline[name]).view(-1) + target = integer_view(self.baseline[name]).view(-1) count = locations.numel() first = int(locations[0]) if int(locations[-1]) - first + 1 == count: diff --git a/nemo_rl/utils/weight_transfer_zmq.py b/nemo_rl/utils/weight_transfer_zmq.py index 624eba3a139..fd3020dde1d 100644 --- a/nemo_rl/utils/weight_transfer_zmq.py +++ b/nemo_rl/utils/weight_transfer_zmq.py @@ -26,7 +26,6 @@ import zmq from nemo_rl.utils.weight_transfer_remote_sparse import ( - SparseDeltaStreamResult, SparsePartitionMode, merge_vllm_refit_metrics, post_vllm_refit_endpoints, @@ -58,6 +57,13 @@ def _json_bytes(value: Mapping[str, Any]) -> bytes: return json.dumps(value, separators=(",", ":"), sort_keys=True).encode() +def _configure_socket(socket: zmq.Socket, high_water_mark: int) -> None: + socket.setsockopt(zmq.LINGER, 0) + socket.setsockopt(zmq.SNDHWM, high_water_mark) + socket.setsockopt(zmq.RCVHWM, high_water_mark) + socket.setsockopt(zmq.TCP_KEEPALIVE, 1) + + class ZmqSparseRefitClient: """One-thread DEALER client with retry-safe payload identifiers.""" @@ -74,16 +80,12 @@ def __init__( self._producer_id = producer_id self._api_key = api_key self._socket = zmq.Context.instance().socket(zmq.DEALER) + _configure_socket(self._socket, 2) self._socket.setsockopt( - zmq.IDENTITY, - f"nrl-{producer_id}-{uuid.uuid4().hex}".encode(), + zmq.IDENTITY, f"nrl-{producer_id}-{uuid.uuid4().hex}".encode() ) - self._socket.setsockopt(zmq.LINGER, 0) self._socket.setsockopt(zmq.IMMEDIATE, 1) - self._socket.setsockopt(zmq.SNDHWM, 2) - self._socket.setsockopt(zmq.RCVHWM, 2) self._socket.setsockopt(zmq.SNDTIMEO, self._timeout_ms) - self._socket.setsockopt(zmq.TCP_KEEPALIVE, 1) self._socket.connect(address) def send_payload( @@ -267,11 +269,6 @@ def _parse_data_message( checksum = str(metadata["checksum"]) if not transfer_id or producer_id < 0 or payload_id < 0: raise ValueError("Invalid ZeroMQ sparse refit payload identity.") - actual = sparse_payload_checksum(body) - if actual != checksum: - raise ValueError( - f"Sparse refit payload checksum mismatch: expected={checksum}, actual={actual}." - ) return identity, (transfer_id, producer_id, payload_id), body, metadata def _run(self) -> None: @@ -287,11 +284,8 @@ def _run(self) -> None: ) pending: dict[Any, tuple[bytes, tuple[str, int, int]]] = {} try: - socket.setsockopt(zmq.LINGER, 0) + _configure_socket(socket, 16) socket.setsockopt(zmq.ROUTER_MANDATORY, 1) - socket.setsockopt(zmq.SNDHWM, 16) - socket.setsockopt(zmq.RCVHWM, 16) - socket.setsockopt(zmq.TCP_KEEPALIVE, 1) socket.bind(self._bind_address) self._endpoint = socket.getsockopt_string(zmq.LAST_ENDPOINT) self._ready.set() @@ -361,7 +355,7 @@ def stream_sparse_delta_payloads_via_zmq( shard_rank: int, shard_count: int, partition: SparsePartitionMode = "chunks", -) -> SparseDeltaStreamResult: +) -> dict[str, int]: addresses = [address.strip() for address in refit_targets if address.strip()] if not addresses: raise ValueError("At least one ZeroMQ sparse refit address is required.") diff --git a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py index 2ab0b40d73c..fa88fb94870 100644 --- a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py +++ b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py @@ -24,9 +24,12 @@ from nemo_rl.utils.timer import Timer from nemo_rl.utils.weight_transfer_remote_sparse import ( - flush_vllm_refit_urls, + G_VLLM_REFIT_FLUSH_PATH, + G_VLLM_REFIT_PREPARE_PATH, merge_vllm_refit_metrics, - prepare_vllm_sparse_refit_urls, + post_vllm_refit_endpoints, + vllm_refit_api_key, + vllm_refit_endpoints, ) from nemo_rl.weight_sync.interfaces import WeightSynchronizer @@ -44,10 +47,12 @@ def validate_vllm_remote_sparse_refit( ) -> str | None: """Validate the optional config and return its internal transport name.""" transport = config.get("refit_transport") - if transport is not None and transport not in _REMOTE_SPARSE_TRANSPORTS: + if transport is None: + return None + if transport not in _REMOTE_SPARSE_TRANSPORTS: raise ValueError(f"Unsupported vLLM refit transport {transport!r}.") vllm_cfg = config["vllm_cfg"] - if transport is not None and ( + if ( colocated or not megatron_enabled or vllm_cfg["precision"] == "fp8" @@ -60,7 +65,7 @@ def validate_vllm_remote_sparse_refit( f"{transport} requires a non-colocated Megatron policy, BF16/FP16 " "vLLM, delta compression, and an unquantized rollout." ) - return None if transport is None else _REMOTE_SPARSE_TRANSPORTS[transport] + return _REMOTE_SPARSE_TRANSPORTS[transport] class VllmRemoteSparseWeightSynchronizer(WeightSynchronizer): @@ -140,11 +145,7 @@ def sync_weights( verification.update( merge_vllm_refit_metrics( {}, - flush_vllm_refit_urls( - self._refit_urls, - api_key_env_var=self._api_key_env_var, - timeout_s=self._request_timeout_s, - ), + self._request_receivers(G_VLLM_REFIT_FLUSH_PATH, {}), maximum=True, candidate_maximum=False, ) @@ -180,9 +181,9 @@ def sync_weights( finally: if not succeeded: with suppress(Exception): - flush_vllm_refit_urls( - self._refit_urls, - api_key_env_var=self._api_key_env_var, + self._request_receivers( + G_VLLM_REFIT_FLUSH_PATH, + {}, timeout_s=min(self._request_timeout_s, 60.0), ) self._baseline_commit_refs = ( @@ -244,6 +245,20 @@ def _run_generation_workers(self, method_name: str, **kwargs: Any) -> list[Any]: ) ) + def _request_receivers( + self, + path: str, + body: dict[str, Any], + *, + timeout_s: float | None = None, + ) -> list[dict[str, Any]]: + return post_vllm_refit_endpoints( + vllm_refit_endpoints(self._refit_urls, path), + body, + api_key=vllm_refit_api_key(self._api_key_env_var), + timeout_s=self._request_timeout_s if timeout_s is None else timeout_s, + ) + @staticmethod def start_baseline(policy: Any, transport: str) -> list[Any]: workers = policy.worker_group @@ -291,11 +306,14 @@ def init_communicator(self) -> None: ray.get(list(self._baseline_init_refs)) ) self._baseline_init_refs.clear() - prepare_vllm_sparse_refit_urls( - self._refit_urls, - state_dict_info, - api_key_env_var=self._api_key_env_var, - timeout_s=self._request_timeout_s, + self._request_receivers( + G_VLLM_REFIT_PREPARE_PATH, + { + "tensors": { + name: [list(shape), str(dtype).removeprefix("torch.")] + for name, (shape, dtype) in state_dict_info.items() + } + }, ) self._stale = False diff --git a/tests/unit/models/generation/test_vllm_sparse_delta.py b/tests/unit/models/generation/test_vllm_sparse_delta.py index d7af593c934..2ea6c851cb9 100644 --- a/tests/unit/models/generation/test_vllm_sparse_delta.py +++ b/tests/unit/models/generation/test_vllm_sparse_delta.py @@ -28,9 +28,17 @@ SparseOperation, decode_sparse_tensor_payload_for_staging, encode_sparse_infos, - integer_dtype_for_element_size, iter_decoded_sparse_payload, ) +from nemo_rl.utils.weight_transfer_sparse_codec import ( + integer_view as _bits, +) + +_PACKED_LOADER_INFOS = [ + ("row_slice", (4, 2), [4, 5, 7], [1.0] * 3), + ("column_slice", (2, 8), [4, 7, 12, 15], [2.0] * 4), + ("split", (6, 2), [2, 5, 8, 11], [3.0, 3.0, 4.0, 4.0]), +] class _NativeLoaderModel: @@ -55,27 +63,15 @@ def load_weights(self, weights): target["identity"].copy_(source) elif name == "weight_scale_inv": target["scale"].copy_(source) - elif name.endswith("self_attn.k_proj.weight"): - target["qkv"][4:6].copy_(source[2:4]) - elif name.endswith("mlp.gate_proj.weight"): - target["merged"][:4].copy_(source[4:8]) - elif name.endswith("mlp.up_proj.weight"): - target["merged"][4:8].copy_(source[4:8]) - elif name.endswith("experts.3.gate_proj.weight"): - target["w13"][1, :4].copy_(source[4:8]) - elif name.endswith("experts.3.down_proj.weight"): - target["w2"][1].copy_(source[:, 4:8]) - elif name.endswith("mixer.in_proj.weight"): - target["mamba"][:2].copy_(source[2:4]) - target["mamba"][2:].copy_(source[6:10]) - elif name.endswith(("mixer.A", "mixer.A_log")): - target["a"].copy_(source) - target["a"].copy_(-torch.exp(target["a"])) - elif name == "transformed": - target["identity"].copy_(source + 1) - elif name == "raise_after_copy": - target["identity"].copy_(source) - raise RuntimeError("loader failed") + elif name == "row_slice": + target[name].copy_(source[2:4]) + elif name == "column_slice": + target[name].copy_(source[:, 4:8]) + elif name == "split": + target[name][:2].copy_(source[1:3]) + target[name][2:].copy_(source[4:6]) + elif name == "exp_transform": + target[name].copy_(-torch.exp(source)) else: continue loaded.add(name) @@ -92,12 +88,6 @@ def _applier(model: Any) -> VllmSparseDeltaApplier: ) -def _bits(values: torch.Tensor) -> torch.Tensor: - return values.contiguous().view( - integer_dtype_for_element_size(values.element_size()) - ) - - def _decode_staged(payload: Any) -> list[Any]: return list( iter_decoded_sparse_payload(decode_sparse_tensor_payload_for_staging(payload)) @@ -117,8 +107,7 @@ def _payload( def _apply_payload(applier: VllmSparseDeltaApplier, payload: Any) -> None: - for item, locations, values in _decode_staged(payload): - applier._apply_decoded_item(item, locations, values) + applier._apply_decoded_items(_decode_staged(payload)) def test_sparse_prewarm_reserves_largest_source_without_loading_weights() -> None: @@ -151,7 +140,7 @@ def test_sparse_prewarm_caches_rank_local_native_loader_skips() -> None: @pytest.mark.vllm -def test_backend_applies_decoded_sparse_payload_files() -> None: +def test_backend_applies_decoded_sparse_payload_sources() -> None: from nemo_rl.models.generation.vllm.vllm_backend import ( VllmInternalWorkerExtension, ) @@ -159,21 +148,18 @@ def test_backend_applies_decoded_sparse_payload_files() -> None: ext = VllmInternalWorkerExtension.__new__(VllmInternalWorkerExtension) applier = MagicMock() applier.update_weights_from_decoded_sparse_payload.return_value = {"ok": True} - applier.update_weights_from_decoded_sparse_payload_files.return_value = {"ok": True} ext._get_sparse_delta_applier = MagicMock(return_value=applier) ext.prepare_sparse_delta_refit_info({"weight": ((8,), torch.float32)}) assert ext.update_weights_from_decoded_sparse_payload(b"payload") == {"ok": True} - assert ext.update_weights_from_decoded_sparse_payload_files("first", "second") == { + assert ext.update_weights_from_decoded_sparse_payload("first", "second") == { "ok": True } - applier.update_weights_from_decoded_sparse_payload.assert_called_once_with( - b"payload" - ) - applier.update_weights_from_decoded_sparse_payload_files.assert_called_once_with( - "first", "second" - ) + assert [ + item.args + for item in applier.update_weights_from_decoded_sparse_payload.call_args_list + ] == [(b"payload",), ("first", "second")] applier.prewarm.assert_called_once_with({"weight": ((8,), torch.float32)}) applier.discover_native_skips.assert_called_once_with( {"weight": ((8,), torch.float32)} @@ -185,29 +171,13 @@ def test_sparse_payload_batches_preserve_order(tmp_path) -> None: applier = _applier(_NativeLoaderModel(identity=torch.zeros(1))) decoded_paths = [tmp_path / f"decoded-{index}.pt" for index in range(3)] decoded_payloads = [ - ( - ( - torch.tensor([index], dtype=torch.int32), - torch.empty(0, dtype=torch.int64), - ), - (torch.tensor([index]),), - [ - { - "name": "weight", - "index": index, - "shape": (1,), - "dtype": "float32", - "operation": "overwrite", - "index_encoding": "range", - "range_start": 0, - "decoded_location_group": 0, - "decoded_location_start": 0, - "decoded_location_end": 1, - "value_group": 0, - "value_start": 0, - "value_end": 1, - } - ], + decode_sparse_tensor_payload_for_staging( + _payload( + f"weight-{index}", + torch.empty(1), + [0], + _bits(torch.tensor([float(index)])), + ) ) for index in range(3) ] @@ -220,20 +190,19 @@ def test_sparse_payload_batches_preserve_order(tmp_path) -> None: result = applier.update_weights_from_decoded_sparse_payload( *(path.read_bytes() for path in decoded_paths) ) - decoded_result = applier.update_weights_from_decoded_sparse_payload_files( + decoded_result = applier.update_weights_from_decoded_sparse_payload( *(str(path) for path in reversed(decoded_paths)) ) - assert [item["index"] for item in decoded_applied] == [ - 0, - 1, - 2, - 2, - 1, - 0, + assert [item["name"] for item in decoded_applied] == [ + "weight-0", + "weight-1", + "weight-2", + "weight-2", + "weight-1", + "weight-0", ] assert result["receiver_deserialize_s"] >= 0.0 - assert result["receiver_scratch_s"] >= 0.0 assert result["receiver_sparse_apply_s"] >= 0.0 assert decoded_result["receiver_deserialize_s"] >= 0.0 @@ -270,38 +239,26 @@ def test_decoded_sparse_payload_converts_compact_locations_for_apply(tmp_path) - result = _applier( _NativeLoaderModel(identity=target) - ).update_weights_from_decoded_sparse_payload_files(str(path)) + ).update_weights_from_decoded_sparse_payload(str(path)) assert torch.equal(target, torch.tensor([0.0, 2.0, 0.0, 0.0, 0.0, 6.0, 0.0, 0.0])) assert result["receiver_sparse_apply_s"] >= 0.0 @pytest.mark.vllm -def test_native_loaders_apply_sparse_overwrite_without_family_plans() -> None: +def test_native_loaders_apply_sparse_views_and_transforms() -> None: targets = { "identity": torch.zeros(4), - "qkv": torch.zeros(8, 2), - "merged": torch.zeros(8, 2), - "w13": torch.zeros(2, 4, 2), - "w2": torch.zeros(2, 2, 4), - "mamba": torch.zeros(6, 2), - "a": torch.tensor([-2.0, -4.0]), + "row_slice": torch.zeros(2, 2), + "column_slice": torch.zeros(2, 4), + "split": torch.zeros(4, 2), + "exp_transform": torch.tensor([-2.0, -4.0]), } infos = [ ("weight", (4,), [1, 2], [1.0, 2.0]), - ("model.layers.0.self_attn.k_proj.weight", (4, 2), [0, 4, 5, 7], [1] * 4), - ("model.layers.0.mlp.gate_proj.weight", (8, 2), [0, 8, 9, 15], [2] * 4), - ("model.layers.0.mlp.up_proj.weight", (8, 2), [8, 15], [3] * 2), - ("model.layers.0.mlp.experts.3.gate_proj.weight", (8, 2), [8, 15], [4] * 2), + *_PACKED_LOADER_INFOS, ( - "model.layers.0.mlp.experts.3.down_proj.weight", - (2, 8), - [3, 4, 7, 12, 15], - [5] * 5, - ), - ("model.layers.0.mixer.in_proj.weight", (10, 2), [0, 4, 7, 12, 19], [6] * 5), - ( - "backbone.layers.0.mixer.A_log", + "exp_transform", (2,), [0, 1], [math.log(3.0), math.log(2.0)], @@ -319,19 +276,26 @@ def test_native_loaders_apply_sparse_overwrite_without_family_plans() -> None: for name, shape, locations, values in infos ] ) - payload[2][7]["verification_samples"] = 2 + payload[2][-1]["verification_samples"] = 2 applier = _applier(_NativeLoaderModel(**targets)) _apply_payload(applier, payload) verification = applier.finish_sparse_delta_refit() assert torch.equal(targets["identity"], torch.tensor([0.0, 1.0, 2.0, 0.0])) - assert targets["qkv"].view(-1)[[8, 9, 11]].tolist() == [1.0, 1.0, 1.0] - assert targets["merged"].view(-1)[[0, 1, 7, 8, 15]].tolist() == [2, 2, 2, 3, 3] - assert targets["w13"].view(-1)[[8, 15]].tolist() == [4, 4] - assert targets["w2"].view(-1)[[8, 11, 12, 15]].tolist() == [5, 5, 5, 5] - assert targets["mamba"].view(-1)[[0, 3, 4, 11]].tolist() == [6, 6, 6, 6] - assert torch.allclose(targets["a"], torch.tensor([-3.0, -2.0])) + assert targets["row_slice"].view(-1).tolist() == [1.0, 1.0, 0.0, 1.0] + assert targets["column_slice"].view(-1).tolist() == [2.0, 0.0, 0.0, 2.0] * 2 + assert targets["split"].view(-1).tolist() == [ + 3.0, + 0.0, + 0.0, + 3.0, + 4.0, + 0.0, + 0.0, + 4.0, + ] + assert torch.allclose(targets["exp_transform"], torch.tensor([-3.0, -2.0])) assert verification["verification_candidates"] == 2 assert verification["verification_samples"] == 2 assert verification["verification_exact_mismatches"] == 0 @@ -340,43 +304,10 @@ def test_native_loaders_apply_sparse_overwrite_without_family_plans() -> None: @pytest.mark.vllm def test_xor_applies_through_packed_native_loaders() -> None: targets = { - "qkv": torch.zeros(8, 2), - "merged": torch.zeros(8, 2), - "w13": torch.zeros(2, 4, 2), - "mamba": torch.zeros(6, 2), + "row_slice": torch.zeros(2, 2), + "column_slice": torch.zeros(2, 4), + "split": torch.zeros(4, 2), } - infos = [ - ( - "model.layers.0.self_attn.k_proj.weight", - (4, 2), - [0, 4, 5, 7], - [1.0] * 4, - ), - ( - "model.layers.0.mlp.gate_proj.weight", - (8, 2), - [0, 8, 9, 15], - [2.0] * 4, - ), - ( - "model.layers.0.mlp.up_proj.weight", - (8, 2), - [8, 15], - [3.0] * 2, - ), - ( - "model.layers.0.mlp.experts.3.gate_proj.weight", - (8, 2), - [8, 15], - [4.0] * 2, - ), - ( - "model.layers.0.mixer.in_proj.weight", - (10, 2), - [0, 4, 7, 12, 19], - [6.0] * 5, - ), - ] payload = encode_sparse_infos( [ ( @@ -388,16 +319,24 @@ def test_xor_applies_through_packed_native_loaders() -> None: ), "xor", ) - for name, shape, locations, values in infos + for name, shape, locations, values in _PACKED_LOADER_INFOS ] ) _apply_payload(_applier(_NativeLoaderModel(**targets)), payload) - assert targets["qkv"].view(-1)[[8, 9, 11]].tolist() == [1.0, 1.0, 1.0] - assert targets["merged"].view(-1)[[0, 1, 7, 8, 15]].tolist() == [2, 2, 2, 3, 3] - assert targets["w13"].view(-1)[[8, 15]].tolist() == [4.0, 4.0] - assert targets["mamba"].view(-1)[[0, 3, 4, 11]].tolist() == [6, 6, 6, 6] + assert targets["row_slice"].view(-1).tolist() == [1.0, 1.0, 0.0, 1.0] + assert targets["column_slice"].view(-1).tolist() == [2.0, 0.0, 0.0, 2.0] * 2 + assert targets["split"].view(-1).tolist() == [ + 3.0, + 0.0, + 0.0, + 3.0, + 4.0, + 0.0, + 0.0, + 4.0, + ] @pytest.mark.vllm @@ -406,7 +345,7 @@ def test_native_loader_explicit_skip_is_accepted() -> None: payload = _payload("skipped", torch.empty(1), [0], _bits(torch.tensor([1.0]))) item, locations, values = _decode_staged(payload)[0] - _applier(model)._apply_decoded_item(item, locations, values) + _applier(model)._apply_decoded_items(((item, locations, values),)) assert torch.equal(model.targets["identity"], torch.zeros(1)) @@ -425,28 +364,17 @@ def test_native_loader_claim_without_copy_fails_closed() -> None: def test_sparse_overwrite_preserves_unselected_transform_inputs() -> None: target = torch.tensor([-2.0, -4.0, -6.0, -8.0]) payload = _payload( - "backbone.layers.0.mixer.A_log", + "exp_transform", target, [1], _bits(torch.tensor([math.log(3.0)])), ) - _apply_payload(_applier(_NativeLoaderModel(a=target)), payload) + _apply_payload(_applier(_NativeLoaderModel(exp_transform=target)), payload) assert torch.allclose(target, torch.tensor([-2.0, -3.0, -6.0, -8.0])) -@pytest.mark.vllm -def test_sparse_overwrite_rolls_back_loader_failure() -> None: - target = torch.tensor([1.0, 2.0]) - payload = _payload("raise_after_copy", target, [1], _bits(torch.tensor([3.0]))) - - with pytest.raises(RuntimeError, match="loader failed"): - _apply_payload(_applier(_NativeLoaderModel(identity=target)), payload) - - assert torch.equal(target, torch.tensor([1.0, 2.0])) - - @pytest.mark.vllm def test_unknown_sparse_operation_fails_closed() -> None: target = torch.zeros(1) @@ -571,9 +499,9 @@ def test_overwrite_casts_absolute_source_values() -> None: ("name", "source", "targets", "error"), [ ( - "backbone.layers.0.mixer.A_log", + "exp_transform", torch.tensor([math.log(2.0)]), - {"a": torch.tensor([-1.0])}, + {"exp_transform": torch.tensor([-1.0])}, "transforms its input", ), ( diff --git a/tests/unit/models/generation/test_vllm_sparse_refit.py b/tests/unit/models/generation/test_vllm_sparse_refit.py index 8e906bf4e01..7f618817baa 100644 --- a/tests/unit/models/generation/test_vllm_sparse_refit.py +++ b/tests/unit/models/generation/test_vllm_sparse_refit.py @@ -162,7 +162,7 @@ def test_sparse_refit_queue_stages_payload_before_batch_is_full( assert not list(tmp_path.iterdir()) assert receiver._worker.llm.collective_rpc.call_args_list[0].args[0] == ( - "update_weights_from_decoded_sparse_payload_files" + "update_weights_from_decoded_sparse_payload" ) @@ -262,7 +262,7 @@ def test_sparse_refit_batch_decodes_once_before_collective_apply( staged_locations: list[int] = [] def collective_rpc(method: str, args: tuple[Any, ...]) -> list[Any]: - assert method == "update_weights_from_decoded_sparse_payload_files" + assert method == "update_weights_from_decoded_sparse_payload" for path in args: staged_locations.extend( int(location) @@ -292,8 +292,8 @@ def collective_rpc(method: str, args: tuple[Any, ...]) -> list[Any]: assert [ item.args[0] for item in receiver._worker.llm.collective_rpc.call_args_list ] == [ - "update_weights_from_decoded_sparse_payload_files", - "update_weights_from_decoded_sparse_payload_files", + "update_weights_from_decoded_sparse_payload", + "update_weights_from_decoded_sparse_payload", ] assert response["payloads"] == 1 assert response["receiver_worker_total_s"] == 1.0 @@ -307,7 +307,7 @@ def test_sparse_refit_batch_drains_workers_before_error_cleanup(tmp_path: Path) def collective_rpc(method: str, args: tuple[Any, ...]) -> list[Any]: nonlocal staged_paths - if method == "update_weights_from_decoded_sparse_payload_files": + if method == "update_weights_from_decoded_sparse_payload": staged_paths = args raise RuntimeError("apply failed") assert method == "synchronize_device" @@ -326,7 +326,7 @@ def collective_rpc(method: str, args: tuple[Any, ...]) -> list[Any]: entry.args[0] for entry in receiver._worker.llm.collective_rpc.call_args_list ] == [ - "update_weights_from_decoded_sparse_payload_files", + "update_weights_from_decoded_sparse_payload", "synchronize_device", ] assert not list(tmp_path.iterdir()) @@ -366,7 +366,7 @@ class AsyncLlm: async def collective_rpc( self, method: str, args: tuple[Any, ...] ) -> list[Any]: - assert method == "update_weights_from_decoded_sparse_payload_files" + assert method == "update_weights_from_decoded_sparse_payload" for path in args: staged_locations.extend( location diff --git a/tests/unit/models/megatron/test_community_import.py b/tests/unit/models/megatron/test_community_import.py index 841e9dfe65d..21d18b03202 100644 --- a/tests/unit/models/megatron/test_community_import.py +++ b/tests/unit/models/megatron/test_community_import.py @@ -30,7 +30,7 @@ def _ensure_package(monkeypatch, name: str) -> ModuleType: if "." in name: parent_name, child_name = name.rsplit(".", 1) parent_module = _ensure_package(monkeypatch, parent_name) - monkeypatch.setattr(parent_module, child_name, module, raising=False) + setattr(parent_module, child_name, module) return module @@ -87,16 +87,14 @@ def _install_runtime_stubs_for_hf_import(monkeypatch): parallel_state = ModuleType("megatron.core.parallel_state") parallel_state.model_parallel_is_initialized = lambda: False monkeypatch.setitem(sys.modules, "megatron.core.parallel_state", parallel_state) - monkeypatch.setattr(core_module, "parallel_state", parallel_state, raising=False) + core_module.parallel_state = parallel_state rerun_state_machine = ModuleType("megatron.core.rerun_state_machine") rerun_state_machine.destroy_rerun_state_machine = lambda: None monkeypatch.setitem( sys.modules, "megatron.core.rerun_state_machine", rerun_state_machine ) - monkeypatch.setattr( - core_module, "rerun_state_machine", rerun_state_machine, raising=False - ) + core_module.rerun_state_machine = rerun_state_machine tensor_parallel = ModuleType("megatron.core.tensor_parallel") tensor_parallel.model_parallel_cuda_manual_seed = lambda seed: None @@ -108,7 +106,7 @@ def _install_runtime_stubs_for_hf_import(monkeypatch): monkeypatch.setitem( sys.modules, "megatron.core.tensor_parallel.random", tensor_parallel_random ) - monkeypatch.setattr(core_module, "tensor_parallel", tensor_parallel, raising=False) + core_module.tensor_parallel = tensor_parallel def test_prefer_nvrx_is_noop_when_strategy_import_fails(monkeypatch): diff --git a/tests/unit/models/policy/test_megatron_remote_sparse_refit.py b/tests/unit/models/policy/test_megatron_remote_sparse_refit.py index d3130d4f4cf..4ea1dc34123 100644 --- a/tests/unit/models/policy/test_megatron_remote_sparse_refit.py +++ b/tests/unit/models/policy/test_megatron_remote_sparse_refit.py @@ -43,12 +43,8 @@ def _detect_parallelism_type(self, _module): return self.parallelism -class _ColumnMapping(_AutoMapping): - pass - - -class _DirectMapping(_AutoMapping): - pass +_ColumnMapping = type("_ColumnMapping", (_AutoMapping,), {}) +_DirectMapping = type("_DirectMapping", (_AutoMapping,), {}) class _GatedMapping(_AutoMapping): @@ -56,12 +52,8 @@ def __init__(self, *, gate, up): super().__init__({"gate": gate, "up": up}) -class _RowMapping(_AutoMapping): - pass - - -class _ReplicatedMapping(_AutoMapping): - pass +_RowMapping = type("_RowMapping", (_AutoMapping,), {}) +_ReplicatedMapping = type("_ReplicatedMapping", (_AutoMapping,), {}) def _install_mapping_types(monkeypatch, remote_refit_type): @@ -73,7 +65,7 @@ def _install_mapping_types(monkeypatch, remote_refit_type): _AutoMapping, { _ColumnMapping: megatron_remote_sparse_refit._COLUMN, - _DirectMapping: megatron_remote_sparse_refit._DIRECT, + _DirectMapping: megatron_remote_sparse_refit._REPLICATED, _GatedMapping: megatron_remote_sparse_refit._GATED, _ReplicatedMapping: megatron_remote_sparse_refit._REPLICATED, _RowMapping: megatron_remote_sparse_refit._ROW, @@ -205,70 +197,58 @@ def test_remote_sparse_globalizes_expert_name(): def test_remote_sparse_projects_bridge_affine_mappings(monkeypatch): _install_mapping_types(monkeypatch, MegatronRemoteSparseRefit) - - def task(mapping, module_name, tensor, global_name): - module = type(module_name, (torch.nn.Module,), {})() - return SimpleNamespace( + cases = ( + ( + _AutoMapping("backbone.layers.0.mixer.D"), + torch.arange(8).view(4, 2), + "decoder.layers.0.mixer.D", + megatron_remote_sparse_refit._COLUMN, + ("backbone.layers.0.mixer.D", (8, 2), 0, 4), + ), + ( + _AutoMapping("backbone.layers.0.mixer.o_proj.weight", parallelism="row"), + torch.arange(8).view(2, 4), + "decoder.layers.0.self_attention.linear_proj.weight", + megatron_remote_sparse_refit._ROW, + ("backbone.layers.0.mixer.o_proj.weight", (2, 8), 1, 4), + ), + ( + _AutoMapping("backbone.layers.0.norm.weight", parallelism="replicated"), + torch.arange(4), + "decoder.layers.0.input_layernorm.weight", + megatron_remote_sparse_refit._REPLICATED, + ("backbone.layers.0.norm.weight", (4,), None, 0), + ), + ) + for mapping, tensor, global_name, kind, expected in cases: + mapping.tp_size = 2 + mapping.tp_rank = 1 + task = SimpleNamespace( mapping=mapping, - megatron_module=module, + megatron_module=torch.nn.Module(), param_weight=tensor, global_param_name=global_name, ) - - column = task( - _AutoMapping("backbone.layers.0.mixer.D"), - "TEColumnParallelLinear", - torch.arange(8).view(4, 2), - "decoder.layers.0.mixer.D", - ) - column.mapping.tp_size = 2 - column.mapping.tp_rank = 1 - row = task( - _AutoMapping("backbone.layers.0.mixer.o_proj.weight", parallelism="row"), - "TERowParallelLinear", - torch.arange(8).view(2, 4), - "decoder.layers.0.self_attention.linear_proj.weight", - ) - row.mapping.tp_size = 2 - row.mapping.tp_rank = 1 - replicated = task( - _AutoMapping("backbone.layers.0.norm.weight", parallelism="replicated"), - "TENorm", - torch.arange(4), - "decoder.layers.0.input_layernorm.weight", - ) - gated = task( - _GatedMapping( + projection = MegatronRemoteSparseRefit._task_local_tensors(task, kind)[0][2] + assert ( + projection.name, + projection.global_shape, + projection.shard_dim, + projection.offset, + ) == expected + + gated = SimpleNamespace( + mapping=_GatedMapping( gate="model.mlp.gate_proj.weight", up="model.mlp.up_proj.weight", ), - "TEColumnParallelLinear", - torch.arange(16).view(8, 2), - "mlp.linear_fc1.weight", + megatron_module=torch.nn.Module(), + param_weight=torch.arange(16).view(8, 2), + global_param_name="mlp.linear_fc1.weight", ) - - column_projection = MegatronRemoteSparseRefit._task_local_tensors( - column, megatron_remote_sparse_refit._COLUMN - )[0][2] - row_projection = MegatronRemoteSparseRefit._task_local_tensors( - row, megatron_remote_sparse_refit._ROW - )[0][2] - replicated_projection = MegatronRemoteSparseRefit._task_local_tensors( - replicated, megatron_remote_sparse_refit._REPLICATED - )[0][2] gated_projections = MegatronRemoteSparseRefit._task_local_tensors( gated, megatron_remote_sparse_refit._GATED ) - - assert column_projection.name == "backbone.layers.0.mixer.D" - assert column_projection.global_shape == (8, 2) - assert column_projection.offsets == (4, 0) - assert row_projection.name == "backbone.layers.0.mixer.o_proj.weight" - assert row_projection.global_shape == (2, 8) - assert row_projection.offsets == (0, 4) - assert replicated_projection.name == "backbone.layers.0.norm.weight" - assert replicated_projection.global_shape == (4,) - assert replicated_projection.offsets == (0,) assert [projection.name for _, _, projection in gated_projections] == [ "model.mlp.gate_proj.weight", "model.mlp.up_proj.weight", diff --git a/tests/unit/utils/test_weight_transfer_remote_sparse.py b/tests/unit/utils/test_weight_transfer_remote_sparse.py index ecbe361c3eb..7853ff08028 100644 --- a/tests/unit/utils/test_weight_transfer_remote_sparse.py +++ b/tests/unit/utils/test_weight_transfer_remote_sparse.py @@ -30,7 +30,6 @@ from nemo_rl.utils.weight_transfer_sparse_codec import ( DeltaCompressionTracker, SparseShardProjection, - _bytewise_diff_mask, encode_sparse_infos, sparse_locations_for_item, ) @@ -93,8 +92,25 @@ def _delta_tracker(encoding: str = "overwrite") -> DeltaCompressionTracker: ) -def test_delta_tracker_commits_only_successful_syncs(monkeypatch) -> None: +def _baseline_names(tensors, *, rank: int, partition: str = "chunks"): + tracker = _BaselineNamesTracker() + weight_transfer_remote_sparse.init_sparse_delta_baseline_from_iterator( + tensors, + delta_tracker=tracker, + shard_rank=rank, + shard_count=2, + transport="zmq", + partition=partition, + ) + return tracker.names + + +@pytest.fixture(autouse=True) +def _in_memory_baseline(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") + + +def test_delta_tracker_commits_only_successful_syncs() -> None: tracker = _delta_tracker() tensor = torch.tensor([1.0, 2.0, 3.0]) tracker.snapshot_baseline([("weight", tensor)]) @@ -107,8 +123,7 @@ def test_delta_tracker_commits_only_successful_syncs(monkeypatch) -> None: assert not tracker.prepare_sparse_delta_payload([("weight", tensor)])[0][2] -def test_delta_tracker_change_summary_is_transactional(monkeypatch) -> None: - monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") +def test_delta_tracker_change_summary_is_transactional() -> None: tracker = _delta_tracker() first = torch.tensor([1.0, 2.0]) second = torch.tensor([3.0, 4.0]) @@ -124,7 +139,6 @@ def test_delta_tracker_change_summary_is_transactional(monkeypatch) -> None: def test_delta_tracker_emits_bounded_verification_budget(monkeypatch) -> None: - monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") monkeypatch.setenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "2") tracker = _delta_tracker() tensor = torch.tensor([1.0, 2.0, 3.0, 4.0]) @@ -140,8 +154,7 @@ def test_delta_tracker_emits_bounded_verification_budget(monkeypatch) -> None: assert (changed, total) == (2, 4) -def test_delta_tracker_commits_exact_source_baseline(monkeypatch) -> None: - monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") +def test_delta_tracker_commits_exact_source_baseline() -> None: tracker = _delta_tracker() tensor = torch.tensor([1.0]) tracker.snapshot_baseline([("weight", tensor)]) @@ -157,8 +170,7 @@ def test_delta_tracker_commits_exact_source_baseline(monkeypatch) -> None: assert torch.equal(tracker.baseline["weight"], tensor) -def test_delta_tracker_xor_encodes_against_baseline(monkeypatch) -> None: - monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") +def test_delta_tracker_xor_encodes_against_baseline() -> None: tracker = _delta_tracker("xor") tensor = torch.tensor([1.0, 2.0, 3.0]) baseline = tensor.clone() @@ -184,8 +196,8 @@ def test_delta_tracker_xor_encodes_against_baseline(monkeypatch) -> None: @pytest.mark.parametrize( ("projection", "expected_locations"), [ - (SparseShardProjection("hf.weight", (2, 4), (0, 2)), [2, 7]), - (SparseShardProjection("hf.weight", (4, 2), (2, 0)), [4, 7]), + (SparseShardProjection("hf.weight", (2, 4), 1, 2), [2, 7]), + (SparseShardProjection("hf.weight", (4, 2), 0, 2), [4, 7]), ], ) def test_delta_tracker_projects_local_shards_to_hf_locations( @@ -193,7 +205,6 @@ def test_delta_tracker_projects_local_shards_to_hf_locations( projection: SparseShardProjection, expected_locations: list[int], ) -> None: - monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") monkeypatch.setenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "2") tracker = DeltaCompressionTracker( {"encoding": "overwrite", "sparse_bucket_size_bytes": 1024}, @@ -238,21 +249,10 @@ def test_sparse_index_encoding_preserves_uint64_locations() -> None: assert torch.equal(decoded, locations) -def test_bytewise_diff_mask_supports_float8() -> None: - baseline = torch.tensor([0x38, 0x7F, 0x00], dtype=torch.uint8).view( - torch.float8_e4m3fn - ) - current = baseline.clone() - current.view(torch.uint8)[1] = 0x7E - - assert _bytewise_diff_mask(current, baseline).tolist() == [False, True, False] - - @pytest.mark.parametrize("encoding", ["xor", "overwrite"]) def test_delta_tracker_encodes_fp8_weight_and_scale_bits( monkeypatch, encoding: str ) -> None: - monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") monkeypatch.setenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "2") tracker = _delta_tracker(encoding) weight = torch.tensor([0x38, 0x40, 0x48], dtype=torch.uint8).view( @@ -324,34 +324,6 @@ def test_refit_http_session_does_not_retry_application_errors() -> None: assert {502, 503, 504} <= set(retry.status_forcelist) -def test_prepare_sparse_refit_urls_serializes_metadata(monkeypatch) -> None: - posts = [] - monkeypatch.setenv("NRL_TEST_REFIT_KEY", "secret") - monkeypatch.setattr( - weight_transfer_remote_sparse, - "post_vllm_refit_endpoints", - lambda *args, **kwargs: posts.append((args, kwargs)) or [{"ok": True}], - ) - - result = weight_transfer_remote_sparse.prepare_vllm_sparse_refit_urls( - [" http://receiver/ "], - {"weight": ((2, 3), torch.bfloat16)}, - api_key_env_var="NRL_TEST_REFIT_KEY", - timeout_s=7.0, - ) - - assert result == [{"ok": True}] - assert posts == [ - ( - ( - ["http://receiver/nemo-rl/refit/prepare"], - {"tensors": {"weight": [[2, 3], "bfloat16"]}}, - ), - {"api_key": "secret", "timeout_s": 7.0}, - ) - ] - - def test_sparse_export_finishes_before_blocked_transfers(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_ZMQ_ENCODE_WORKERS", "1") monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "1") @@ -453,62 +425,31 @@ def test_sparse_baseline_snapshots_only_owned_export_chunks( monkeypatch, capsys ) -> None: monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "4") - tracker = _BaselineNamesTracker() - weight_transfer_remote_sparse.init_sparse_delta_baseline_from_iterator( - [(f"weight-{index}", torch.tensor([float(index)])) for index in range(4)], - delta_tracker=tracker, - shard_rank=1, - shard_count=2, - transport="zmq", - ) + tensors = [(f"weight-{index}", torch.tensor([float(index)])) for index in range(4)] - assert tracker.names == ["weight-1", "weight-3"] + assert _baseline_names(tensors, rank=1) == ["weight-1", "weight-3"] assert "chunks=4" in capsys.readouterr().out - - tracker = _BaselineNamesTracker() - weight_transfer_remote_sparse.init_sparse_delta_baseline_from_iterator( - [(f"weight-{index}", torch.tensor([float(index)])) for index in range(4)], - delta_tracker=tracker, - shard_rank=1, - shard_count=2, - transport="zmq", - partition="none", - ) - assert tracker.names == [f"weight-{index}" for index in range(4)] + assert _baseline_names(tensors, rank=1, partition="none") == [ + f"weight-{index}" for index in range(4) + ] def test_sparse_name_partition_is_stable_for_filtered_exports(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "4") tensors = [(f"weight-{index}", torch.tensor([float(index)])) for index in range(8)] - owners = [] - for rank in range(2): - tracker = _BaselineNamesTracker() - weight_transfer_remote_sparse.init_sparse_delta_baseline_from_iterator( - tensors, - delta_tracker=tracker, - shard_rank=rank, - shard_count=2, - transport="zmq", - partition="names", - ) - owners.append(set(tracker.names)) + owners = [ + set(_baseline_names(tensors, rank=rank, partition="names")) for rank in range(2) + ] assert owners[0].isdisjoint(owners[1]) assert owners[0] | owners[1] == {name for name, _tensor in tensors} filtered = tensors[::2] for rank in range(2): - tracker = _BaselineNamesTracker() - weight_transfer_remote_sparse.init_sparse_delta_baseline_from_iterator( - filtered, - delta_tracker=tracker, - shard_rank=rank, - shard_count=2, - transport="zmq", - partition="names", - ) - assert set(tracker.names) == owners[rank] & {name for name, _tensor in filtered} + assert set(_baseline_names(filtered, rank=rank, partition="names")) == owners[ + rank + ] & {name for name, _tensor in filtered} def test_s3_manifest_transport_validates_configuration(monkeypatch) -> None: @@ -760,7 +701,6 @@ def frames(kind: bytes = b"DATA", **updates: object) -> list[bytes]: (frames(protocol="other"), ValueError, "protocol"), (frames(api_key="wrong"), PermissionError, "authentication"), (frames(transfer_id=""), ValueError, "identity"), - (frames(checksum="wrong"), ValueError, "checksum mismatch"), ): with pytest.raises(error, match=match): server._parse_data_message(message) @@ -775,8 +715,13 @@ def do_POST(self): received.append( ({key.lower(): value for key, value in self.headers.items()}, body) ) - response = json.dumps({"ok": True, "receiver_total_s": 0.25}).encode() - self.send_response(200) + ok = self.headers[G_VLLM_REFIT_CHECKSUM_HEADER] == sparse_payload_checksum( + body + ) + response = json.dumps( + {"ok": ok, "receiver_total_s": 0.25, "error": "checksum mismatch"} + ).encode() + self.send_response(200 if ok else 500) self.send_header("content-type", "application/json") self.send_header("content-length", str(len(response))) self.end_headers() 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 04ffe5962c5..ebc55cbb4f1 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 @@ -21,6 +21,20 @@ validate_vllm_remote_sparse_refit, ) +_MODULE = "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer" + + +@pytest.fixture +def mock_ray(): + with patch(f"{_MODULE}.ray") as value: + yield value + + +@pytest.fixture +def post(): + with patch(f"{_MODULE}.post_vllm_refit_endpoints") as value: + yield value + def _remote_sparse_sync( mock_ray: MagicMock, @@ -50,10 +64,7 @@ def _remote_sparse_sync( mock_ray.get.side_effect = get_results sync = VllmRemoteSparseWeightSynchronizer(policy, generation, transport=transport) - with patch( - "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer." - "prepare_vllm_sparse_refit_urls" - ): + with patch(f"{_MODULE}.post_vllm_refit_endpoints"): sync.init_communicator() return sync, policy, generation @@ -101,12 +112,7 @@ def test_validate_remote_sparse_refit_rejects_unsupported_scope(change, kwargs): class TestVllmRemoteSparseWeightSynchronizer: - @patch( - "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer." - "prepare_vllm_sparse_refit_urls" - ) - @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") - def test_init_communicator_joins_prelaunched_baseline(self, mock_ray, prepare): + def test_init_communicator_joins_prelaunched_baseline(self, mock_ray, post): baseline_ref = MagicMock() policy = MagicMock() generation = MagicMock() @@ -126,10 +132,10 @@ def test_init_communicator_joins_prelaunched_baseline(self, mock_ray, prepare): policy.worker_group.run_all_workers_multiple_data.assert_not_called() mock_ray.get.assert_any_call([baseline_ref]) - prepare.assert_called_once_with( - ["http://receiver"], - {"weight": ((8,), "float32")}, - api_key_env_var=None, + post.assert_called_once_with( + ["http://receiver/nemo-rl/refit/prepare"], + {"tensors": {"weight": [[8], "float32"]}}, + api_key=None, timeout_s=600.0, ) assert sync._baseline_init_refs == [] @@ -143,7 +149,6 @@ def test_merge_refit_info_rejects_conflicting_metadata(self): ] ) - @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") def test_init_communicator_requires_receiver_endpoints(self, mock_ray): policy = MagicMock() policy.worker_group.workers = [object()] @@ -157,7 +162,6 @@ def test_init_communicator_requires_receiver_endpoints(self, mock_ray): with pytest.raises(ValueError, match="endpoints are missing"): sync.init_communicator() - @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") def test_shutdown_cancels_pending_work_and_stops_zmq(self, mock_ray): policy = MagicMock() generation = MagicMock() @@ -187,8 +191,7 @@ def test_shutdown_cancels_pending_work_and_stops_zmq(self, mock_ray): assert sync._refit_urls == [] assert sync._targets == [] - @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") - def test_fails_before_transfer_when_kv_cache_invalidation_fails(self, _mock_ray): + def test_fails_before_transfer_when_kv_cache_invalidation_fails(self, mock_ray): policy = MagicMock() generation = MagicMock() generation.invalidate_kv_cache.return_value = False @@ -198,19 +201,15 @@ def test_fails_before_transfer_when_kv_cache_invalidation_fails(self, _mock_ray) sync.sync_weights() policy.worker_group.run_all_workers_multiple_data.assert_not_called() - @patch( - "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.flush_vllm_refit_urls" - ) - @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") def test_initializes_streams_commits_and_updates_baseline( - self, mock_ray, flush, capsys + self, mock_ray, post, capsys ): sync, policy, generation = _remote_sparse_sync( mock_ray, "zmq", [{"payloads": 3, "changed_elements": 3, "total_elements": 100}], ) - flush.return_value = [ + post.return_value = [ { "verification_candidates": 4, "verification_samples": 4, @@ -238,8 +237,11 @@ def test_initializes_streams_commits_and_updates_baseline( run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], refit_urls=["http://receiver"], ) - flush.assert_called_once_with( - ["http://receiver"], api_key_env_var=None, timeout_s=600.0 + post.assert_called_once_with( + ["http://receiver/nemo-rl/refit/flush"], + {}, + api_key=None, + timeout_s=600.0, ) policy.worker_group.run_all_workers_single_data.assert_called_once_with( "finish_remote_sparse_delta_sync", succeeded=True @@ -258,17 +260,13 @@ def test_initializes_streams_commits_and_updates_baseline( assert metrics["transfer/payloads"] == 3.0 assert not sync.is_stale - @patch( - "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.flush_vllm_refit_urls" - ) - @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") - def test_sample_mismatch_does_not_commit_baseline(self, mock_ray, flush): + def test_sample_mismatch_does_not_commit_baseline(self, mock_ray, post): sync, policy, _ = _remote_sparse_sync( mock_ray, "zmq", [{"payloads": 3, "changed_elements": 3, "total_elements": 100}], ) - flush.return_value = [ + post.return_value = [ { "verification_samples": 4, "verification_mismatches": 1, @@ -284,13 +282,7 @@ def test_sample_mismatch_does_not_commit_baseline(self, mock_ray, flush): "finish_remote_sparse_delta_sync", succeeded=False ) - @patch( - "nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.flush_vllm_refit_urls" - ) - @patch("nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer.ray") - def test_failure_drains_receivers_without_committing_baseline( - self, mock_ray, flush - ): + def test_failure_drains_receivers_without_committing_baseline(self, mock_ray, post): sync, policy, _ = _remote_sparse_sync( mock_ray, "s3", RuntimeError("stream failed") ) @@ -298,8 +290,11 @@ def test_failure_drains_receivers_without_committing_baseline( with pytest.raises(RuntimeError, match="stream failed"): sync.sync_weights() - flush.assert_called_once_with( - ["http://receiver"], api_key_env_var=None, timeout_s=60.0 + post.assert_called_once_with( + ["http://receiver/nemo-rl/refit/flush"], + {}, + api_key=None, + timeout_s=60.0, ) policy.worker_group.run_all_workers_single_data.assert_called_once_with( "finish_remote_sparse_delta_sync", succeeded=False diff --git a/tools/refit_bandwidth_calculator.py b/tools/refit_bandwidth_calculator.py index 78f3b9772cd..2dbeb34fbc9 100644 --- a/tools/refit_bandwidth_calculator.py +++ b/tools/refit_bandwidth_calculator.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Compare measured sparse refit with NCCL projected from H100 IB to Ethernet.""" +"""Compare measured zstd sparse refit with NCCL projected to Ethernet.""" import argparse import json @@ -22,33 +22,30 @@ from typing import Literal Transport = Literal["s3", "zmq"] -Compression = Literal["raw", "zstd"] _REFERENCE_IB_GBPS = 400.0 _DENSITIES = (3.0, 5.0) _SPARSE_SIZE_RANGE_GB = (63.2, 1121.0) -# Each pair is (fixed seconds, seconds per 1,000 GB) at 3% and 5%. +# Each pair is (fixed seconds, seconds per 1,000 GB) at 3% and 5% changed. _SPARSE_LATENCY_FITS: dict[ - tuple[Transport, Compression], tuple[tuple[float, float], tuple[float, float]] + Transport, tuple[tuple[float, float], tuple[float, float]] ] = { - ("s3", "raw"): ((2.370, 116.138), (2.646, 183.431)), - ("s3", "zstd"): ((0.677, 84.672), (3.741, 145.597)), - ("zmq", "raw"): ((0.000, 264.753), (0.000, 428.008)), - ("zmq", "zstd"): ((5.813, 73.257), (0.000, 169.991)), + "s3": ((2.271898, 90.539029), (9.132526, 164.284601)), + "zmq": ((5.290113, 78.016171), (10.586558, 124.416925)), } +# Unique compressed wire GB per 1,000 GB of indexed BF16 weights at 3% and 5%. +_SPARSE_WIRE_GB_PER_TB = (22.430807, 37.362273) _NCCL_ANCHORS = ( (63.2, 0.84, 1.60), (247.2, 1.46, 1.74), (470.2, 2.31, 2.73), (1342.0, 3.27, 3.46), ) -_WIRE_MULTIPLIER: dict[Compression, float] = {"raw": 2.0, "zstd": 0.747} @dataclass(frozen=True) class Estimate: transport: Transport - compression: Compression model_size_gb: float changed_pct: float sparse_seconds: float @@ -68,18 +65,26 @@ def predict_sparse_seconds( changed_pct: float, *, transport: Transport, - compression: Compression, ) -> float: """Evaluate the campaign fit at any positive changed density.""" if model_size_gb <= 0 or changed_pct <= 0: raise ValueError("model_size_gb and changed_pct must be positive") size_tb = model_size_gb / 1000.0 - (a3, b3), (a5, b5) = _SPARSE_LATENCY_FITS[(transport, compression)] + (a3, b3), (a5, b5) = _SPARSE_LATENCY_FITS[transport] latency_3, latency_5 = a3 + b3 * size_tb, a5 + b5 * size_tb exponent = math.log(latency_5 / latency_3) / math.log(5.0 / 3.0) return latency_3 * (changed_pct / 3.0) ** exponent +def predict_sparse_wire_gb(model_size_gb: float, changed_pct: float) -> float: + """Scale the measured zstd wire fit to model size and changed density.""" + if model_size_gb <= 0 or changed_pct <= 0: + raise ValueError("model_size_gb and changed_pct must be positive") + wire_3, wire_5 = _SPARSE_WIRE_GB_PER_TB + exponent = math.log(wire_5 / wire_3) / math.log(5.0 / 3.0) + return model_size_gb / 1000.0 * wire_3 * (changed_pct / 3.0) ** exponent + + def _nccl_reference(model_size_gb: float) -> tuple[float, float]: sizes = tuple(anchor[0] for anchor in _NCCL_ANCHORS) index = min(max(bisect_right(sizes, model_size_gb) - 1, 0), len(sizes) - 2) @@ -95,7 +100,6 @@ def estimate( model_size_gb: float, changed_pct: float, transport: Transport, - compression: Compression = "zstd", candidate_ethernet_gbps: float | None = None, ) -> Estimate: """Estimate sparse latency and the NCCL-over-Ethernet crossover.""" @@ -105,7 +109,6 @@ def estimate( model_size_gb, changed_pct, transport=transport, - compression=compression, ) nccl_low, nccl_high = _nccl_reference(model_size_gb) crossover_low = _REFERENCE_IB_GBPS * nccl_low / sparse_seconds @@ -129,11 +132,10 @@ def estimate( return Estimate( transport, - compression, model_size_gb, changed_pct, sparse_seconds, - model_size_gb * changed_pct / 100 * _WIRE_MULTIPLIER[compression], + predict_sparse_wire_gb(model_size_gb, changed_pct), nccl_low, nccl_high, candidate_ethernet_gbps, @@ -160,7 +162,7 @@ def _print_results(results: list[Estimate]) -> None: first = results[0] print( f"Model: {first.model_size_gb:g} GB indexed BF16; " - f"changed: {first.changed_pct:g}%; compression: {first.compression}" + f"changed: {first.changed_pct:g}%; compression: zstd" ) print( "Measured NCCL on 400 Gbps/rank H100 IB: " @@ -204,7 +206,6 @@ def main() -> None: "--changed-pct", "--sparsity-pct", type=_positive, required=True ) parser.add_argument("--transport", choices=("all", "s3", "zmq"), default="all") - parser.add_argument("--compression", choices=("raw", "zstd"), default="zstd") parser.add_argument("--candidate-ethernet-gbps", type=_positive) parser.add_argument("--json", action="store_true") args = parser.parse_args() @@ -216,7 +217,6 @@ def main() -> None: model_size_gb=args.model_size_gb, changed_pct=args.changed_pct, transport=transport, - compression=args.compression, candidate_ethernet_gbps=args.candidate_ethernet_gbps, ) for transport in transports From caf156fa93cd425eddfd614180072cf5f3b65b19 Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Wed, 15 Jul 2026 15:45:24 -0700 Subject: [PATCH 12/18] Remove mcore local baseline and clean up Signed-off-by: Hollow Man --- docs/design-docs/sparse-delta-refit.md | 325 +++++++------ examples/configs/grpo_math_1B.yaml | 2 +- nemo_rl/models/generation/vllm/config.py | 8 +- .../models/generation/vllm/vllm_backend.py | 15 +- .../generation/vllm/vllm_sparse_delta.py | 127 +++-- .../generation/vllm/vllm_sparse_refit.py | 121 +++-- nemo_rl/models/generation/vllm/vllm_worker.py | 6 + .../policy/workers/megatron_policy_worker.py | 7 +- .../workers/megatron_remote_sparse_refit.py | 455 +----------------- nemo_rl/utils/weight_transfer_http.py | 141 ++++++ nemo_rl/utils/weight_transfer_sparse_codec.py | 153 +----- ...te_sparse.py => weight_transfer_stream.py} | 345 ++++++------- nemo_rl/utils/weight_transfer_zmq.py | 377 +++++++++++---- .../vllm_remote_sparse_weight_synchronizer.py | 69 ++- pyrefly.toml | 3 +- .../models/generation/test_vllm_backend.py | 2 +- .../generation/test_vllm_sparse_delta.py | 101 ++-- .../generation/test_vllm_sparse_refit.py | 93 +++- .../test_megatron_remote_sparse_refit.py | 441 +++-------------- .../unit/reference_configs/grpo_math_1B.yaml | 2 +- ...arse.py => test_weight_transfer_stream.py} | 333 ++++++++----- ..._vllm_remote_sparse_weight_synchronizer.py | 78 ++- tools/refit_bandwidth_calculator.py | 66 ++- 23 files changed, 1519 insertions(+), 1751 deletions(-) create mode 100644 nemo_rl/utils/weight_transfer_http.py rename nemo_rl/utils/{weight_transfer_remote_sparse.py => weight_transfer_stream.py} (71%) rename tests/unit/utils/{test_weight_transfer_remote_sparse.py => test_weight_transfer_stream.py} (74%) diff --git a/docs/design-docs/sparse-delta-refit.md b/docs/design-docs/sparse-delta-refit.md index b84ac451b4b..707435fe227 100644 --- a/docs/design-docs/sparse-delta-refit.md +++ b/docs/design-docs/sparse-delta-refit.md @@ -1,11 +1,11 @@ # Remote Sparse-Delta vLLM Refit Remote sparse refit updates non-colocated vLLM workers without transferring a -full checkpoint after every optimizer step. Megatron workers compare uniquely -owned MCore shards against a policy-local CPU baseline. Affine Bridge mappings -emit changed values in canonical Hugging Face (HF) coordinates; only residual -conversion tasks use `export_hf_weights()`. S3 and ZeroMQ share the codec, -pipeline, receiver, native-loader apply path, and commit protocol. +full checkpoint after every optimizer step. Megatron Bridge exports canonical +Hugging Face (HF) tensors, and the policy workers compare their assigned export +chunks against one distributed canonical CPU baseline. S3 and ZeroMQ share the +codec, streaming pipeline, receiver, native-loader apply path, and commit +protocol. The feature is opt-in and does not change existing NCCL, CUDA IPC, or packed refit behavior. @@ -33,14 +33,9 @@ applied and the global flush completes. ```mermaid flowchart LR subgraph P["Megatron policy cluster"] - T["Bridge conversion-task metadata"] - L["All unique MCore shards"] - B["Changed residual Bridge export"] - C["Local source and sharded HF baselines"] + B["Full canonical Bridge export"] + C["Chunk-sharded canonical HF baseline"] E["Compare, encode, and compress"] - T --> L - T --> B - L --> C B --> C C --> E end @@ -50,7 +45,7 @@ flowchart LR subgraph G["vLLM generation cluster"] H["HTTP receiver"] - Q["Eager node staging and bounded FIFO apply queue"] + Q["Compact node staging and bounded FIFO apply queue"] A["Reusable dense scratch and native load_weights()"] H --> Q --> A end @@ -65,9 +60,10 @@ apply engine, and commit protocol.* | Responsibility | Implementation | |---|---| | Coordinate one transfer and commit | [`vllm_remote_sparse_weight_synchronizer.py`](../../nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py) | -| Adapt Megatron workers and assign ownership | [`megatron_remote_sparse_refit.py`](../../nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py) | +| Export canonical tensors and own the source tracker | [`megatron_remote_sparse_refit.py`](../../nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py) | | Track baselines and encode deltas | [`weight_transfer_sparse_codec.py`](../../nemo_rl/utils/weight_transfer_sparse_codec.py) | -| Run the shared pipeline and S3 transport | [`weight_transfer_remote_sparse.py`](../../nemo_rl/utils/weight_transfer_remote_sparse.py) | +| Run the shared pipeline and S3 transport | [`weight_transfer_stream.py`](../../nemo_rl/utils/weight_transfer_stream.py) | +| Share HTTP control-plane utilities | [`weight_transfer_http.py`](../../nemo_rl/utils/weight_transfer_http.py) | | Run the ZeroMQ transport and relay | [`weight_transfer_zmq.py`](../../nemo_rl/utils/weight_transfer_zmq.py) | | Queue receiver work and expose endpoints | [`vllm_sparse_refit.py`](../../nemo_rl/models/generation/vllm/vllm_sparse_refit.py) | | Apply canonical updates through native loaders | [`vllm_sparse_delta.py`](../../nemo_rl/models/generation/vllm/vllm_sparse_delta.py) | @@ -83,43 +79,28 @@ endpoints and then joins the prelaunched baseline before setup returns. The first rollout therefore does not enter a redundant weight sync or race an unfinished snapshot with policy training. -Conversion tasks are split into deterministic paths without changing Megatron -Bridge: - -- Every conversion task keeps its source baseline in MCore layout. A stable - name hash assigns replicated dense tensors across their combined DP/CP and TP - replicas, and expert tensors across their expert-DP replicas. TP, PP, EP, and - ETP still contribute their unique shards. This keeps one source copy while - sharing baseline scans and uploads across equivalent ranks. -- Exact direct, column, row, replicated, and gated mappings produce canonical - HF-coordinate deltas without Bridge export. The decision uses the resolved - Bridge mapping type, not a model name or parameter suffix. -- Tasks with a transform also keep a canonical HF baseline, sharded across - workers by a stable hash of the HF name. Stable ownership is required because - later refits export only changed residual tasks and therefore produce a - different chunk sequence. - -Changed flat locations on an affine mapping are projected with its shard -dimension, TP or ETP rank, and EP-global expert number. Attention output -projections, Mamba affine weights, norms, routers, shared experts, and other -exact mappings use this same model-family-agnostic path. - -Transformed, grouped, padded, tied, adapter, custom-postprocessed, custom FSDP, -and FP8-parameter tasks use the residual path. One integer flag per conversion -task is reduced across the policy world, and Bridge exports only globally -changed tasks when their dependencies are known. If one member of a grouped -export changes, the complete group is exported. A custom Bridge postprocessor -can have undeclared cross-task dependencies, so it retains the full residual -set whenever any residual task changes. Compound QKV, Mamba packing, -permutations, fused exports, and padded or tied embeddings therefore preserve -Bridge semantics. - -The local baselines contain one copy of each unique source element across the -policy workers, rather than a full HF copy per DP replica. Residual tasks have -one additional canonical HF copy distributed across workers. Baselines use -file-backed `torch.from_file` tensors by default; -`NRL_REFIT_BASELINE_IN_MEMORY=1` keeps them in RAM. Local snapshotting and the -residual Bridge export run concurrently. +Every policy rank participates in Megatron Bridge's normal +`export_hf_weights()` conversion. The bounded output chunks are assigned by +`chunk_index % shard_count`, so only one producer snapshots, compares, encodes, +and sends each canonical chunk. Across all producers, persistent baseline +storage is approximately one full HF checkpoint, not one copy per producer. +Skipped chunks are still produced transiently because Bridge collectives and +conversion must run on every participating rank. + +Ownership is deterministic for a fixed canonical export order, tensor shapes, +export chunk limit, and producer count. It is recomputed rather than stored in a +manifest, so changing any of those inputs can move a chunk to another producer; +within one transfer, every chunk still has exactly one owner. + +This deliberately keeps one representation and one tracker. There is no MCore +policy-local baseline, conversion-task detector, affine projection, stable-name +residual partition, or model-specific mapping logic in NeMo RL. QKV, MoE, +Mamba, padded or tied weights, grouped exports, adapters, and custom Bridge +postprocessing all follow Bridge's canonical export semantics. + +Baselines use file-backed `torch.from_file` tensors by default; +`NRL_REFIT_BASELINE_IN_MEMORY=1` keeps them in RAM. File backing reduces +anonymous resident-memory pressure but does not reduce logical baseline bytes. Baseline initialization also returns each canonical tensor's name, shape, and dtype. The synchronizer merges that metadata and asks every vLLM worker to @@ -134,73 +115,88 @@ baseline from an arbitrary training checkpoint. ### Compare and encode -Every uniquely owned local tensor is copied to CPU and compared bytewise through -an integer view with the same element width. The payload contains only changed -locations and values. The comparison remains proportional to model size because -every unique source element is copied and scanned; Adam can therefore make it -CPU-bound even when the wire payload is sparse. - -Policy-local comparison removes full-tensor TP/EP gathers and PP broadcasts for -directly projectable weights. Let `P` be directly projectable bytes, `R` -residual source bytes, `R_changed` the residual HF tensors selected by task -flags, and `s` the element change fraction. The leading work is approximately: +Each producer compares only its assigned canonical export chunks bytewise +through an integer view with the same element width. The payload contains only +changed locations and values. Let `H` be the full canonical checkpoint bytes, +`W` the producer count, and `s` the changed element fraction. Aggregate +persistent baseline, comparison, and wire volumes are approximately: ```text -old: Bridge(P + R) + HF D2H/scan(P + R) + wire(s(P + R)) -new: local D2H/scan(P + R) + Bridge(R_changed) + HF scan(R_changed) - + wire(s(P + R)) + one O(number_of_tasks) flag all-reduce +per-producer baseline and compare ~= H / W +aggregate baseline and compare = H +aggregate wire = codec_metadata + changed_indices + sH ``` -When Adam changes at least one element in every residual tensor, `R_changed` -approaches `R`. The gain then comes from removing Bridge work for `P`, not from -task filtering. It helps less when local D2H or CPU scanning is the bottleneck, -most bytes use custom transformations, or `s` is high. The reported changed -percentage is computed from unique policy-local source elements, so the residual -detector does not inflate it with a second HF scan. +The unavoidable leading cost is still `Bridge_export(H)`: every rank must +participate in TP/EP gathers, PP communication, and conversion before chunk +ownership can discard unassigned output. Sparse refit therefore makes wire and +receiver work proportional to `s`, but not Bridge export communication. The +reported changed percentage is computed once from the assigned canonical chunks +and aggregated across producers. `DeltaCompressionTracker` finds changed flat locations through the equal-width -integer view. Both encodings are dtype-blind and preserve FP8 bit patterns in -the codec, although end-to-end FP8 rollout refit is outside the supported scope. -`overwrite` carries absolute new bits, is idempotent, and is recommended. `xor` -carries new bits XOR baseline bits and may compress better, but requires exact -same-dtype state and exactly-once application. Selecting XOR creates -mixed-operation payloads: direct bitwise-compatible policy shards use XOR, -while Bridge residuals and the full-HF compatibility path use overwrite. A -payload batch can contain both operations. The receiver fails closed on a -transform, dtype cast, or overlapping XOR destination. Use overwrite when any -native loader cannot preserve those XOR requirements. The pending source -baseline always records the exact new source bits. +integer view. During receiver prewarm, the generic apply context observes each +native loader's tensor copies. A name remains XOR-compatible only when every +source is a same-dtype view of the canonical scratch storage and destination +copies do not overlap. The receiver returns the union of incompatible names; +the producer emits XOR for all other names and absolute overwrite values for +that opaque set. This runtime classification contains no model-family or tensor +layout rules. The pending source baseline always records the exact new source +bits regardless of the wire operation. The producer pulls bounded export chunks and compares them in parallel. A separate bounded stage coalesces encoded chunks up to `sparse_bucket_size_bytes`, serializes them, and applies zstd level 1 before the -transport executor. Keeping the 256 MiB compare chunk smaller than the -recommended 1 GiB wire bucket preserves D2H/scan parallelism while reducing -object and manifest count. Payload N can transfer while later chunks are -compared and encoded. Source baselines do not commit until the complete transfer -and receiver flush succeed. Pipeline errors cancel outstanding local futures -and propagate to the synchronizer. +transport executor. The measured S3 default keeps 64 MiB compare chunks smaller +than the recommended 512 MiB wire bucket, preserving D2H/scan parallelism while +reducing object and manifest count. ZeroMQ retains 256 MiB export chunks and +uses the same 512 MiB wire-bucket default, producing 482 payloads across the +current 32-producer topology while tree delivery overlaps later export chunks. + +Payload N can transfer while later chunks are compared and encoded. Source +baselines do not commit until the complete transfer and all transport and +receiver completion barriers succeed. Pipeline errors cancel outstanding local +futures and propagate to the synchronizer. ### Transport and apply | Property | S3 | ZeroMQ | |---|---|---| | Value plane | AWS CRT multipart `PUT_OBJECT` | DEALER to ROUTER relay | -| Notification | HTTP object manifest | Relay HTTP fanout | +| Producer completion | HTTP object manifest accepted by every receiver | Relay registers payload and returns a staged ACK | | Retry identity | Object key + checksum | Transfer + producer + payload IDs + checksum | -| Lifetime | Delete after all receivers reply | No persistent object | +| Completion barrier | Receiver flush | Relay-tree flush, then receiver flush | +| Cleanup | Delete after receiver replies; retry leftovers at stream end | Close producer sockets after sends; seal the transfer at flush | S3 uses 64 MiB multipart parts, a 2 GiB CRT client memory limit, and a 10 Gbps -throughput target. ZeroMQ assigns each producer to one inference-cluster relay; -that relay fans compressed bytes to every generation replica and avoids -duplicate cross-cluster traffic. Both transports use the same receiver -endpoints and checksum validation. +throughput target. ZeroMQ assigns each producer to one inference-cluster relay. +That root applies the payload locally and forwards it through a balanced binary +relay tree, so each payload crosses the inter-cluster boundary once rather than +once per generation replica. Both transports use the same receiver endpoints +and checksum validation. + +Each ZeroMQ relay validates and deduplicates `(transfer_id, producer_id, +payload_id)`, submits its local apply and at most two child forwards to bounded +executors, and acknowledges once that work is registered. The producer can then +export and send later payloads while earlier tree delivery continues. After +every producer finishes, the synchronizer calls `/nemo-rl/refit/zmq-flush` on +all relays and checks that their staged payload count matches the +producer-reported payload total. Only after all tree and apply futures succeed +does it call the normal receiver flush. Tree delivery therefore overlaps the +producer stream, but its remaining tail is still on the final critical path. + +The generic streamer accepts a `SparseRefitTransport` with `send()` and +`cleanup()` methods. It owns export, compare, encode, serialization, +backpressure, timing, and transfer concurrency. S3 and ZeroMQ only construct +their transport state and call that streamer. Cleanup runs on each transfer +worker after all sends finish, so failed S3 deletion retries and ZeroMQ socket +closure occur on the same worker that created the resource. The receiver deduplicates payload identities and applies bounded batches on one FIFO worker thread. Each generation replica downloads a transport payload once. -When its vLLM ranks share a node, decode and flat-file staging under `/dev/shm` -begin as soon as each payload arrives, without waiting for the batch to fill. -Staging futures feed the serial collective apply worker, so download, +When its vLLM ranks share a node, the compact serialized payload is written +directly under `/dev/shm` as soon as it arrives, without waiting for the batch +to fill. Staging futures feed the serial collective apply worker, so download, decompression, staging, and earlier GPU applies can overlap. Queue depth limits submitted work to 32 batches by default; with batches of eight, that is roughly 256 payloads plus the current partial batch. @@ -209,11 +205,10 @@ Locations use `int32` unless a single canonical tensor exceeds the signed 32-bit index range; values remain grouped by dtype. For shared-node workers the collective RPC passes only staged file paths, and each rank uses `torch.load(..., mmap=True)`. The mmap is not a second baseline: it lets ranks -share the staged file's page cache instead of materializing independent CPU -copies. The format flattens locations into one `int32` and one `int64` tensor; -it does not serialize one tensor object per model parameter. When ranks do not -share a node, the receiver still decodes once and sends the same flat -representation through one collective RPC. +share the compact payload's page cache instead of materializing independent CPU +copies. Each worker decodes locations lazily while feeding the native loader. +When ranks do not share a node, the receiver sends the same serialized bytes +through one collective RPC and each worker decodes its copy. There is deliberately no TP/EP source plan. Every vLLM worker sees canonical sparse entries, scatters them into its reusable dense source buffer, and calls @@ -225,14 +220,16 @@ internals. Measure that cost on the target TP/EP topology; H2D no longer scales only with the worker-local sparse subset. During the untimed metadata prewarm, one no-op native-loader pass records names -that issue no model-storage copy on that fixed rank. Later refits skip scratch -construction for those explicit pipeline, expert, or MTP skips. The cache -contains names only; it stores no placement offsets, tensor routes, or -model-family rules. +that issue no model-storage copy on a fixed rank and identifies loaders for +which XOR cannot preserve copy semantics. Later refits skip rank-local names +that the loader explicitly omitted and use overwrite for the incompatible +union. The cache contains names only; it stores no placement offsets, tensor +routes, or model-family rules. -The final `/nemo-rl/refit/flush` drains every batch, synchronizes CUDA, and -checks optional delta samples. Only then does the source commit exact pending -baseline bits in background CPU threads. +For ZeroMQ, the relay flush first drains every staged tree delivery. The final +`/nemo-rl/refit/flush` then drains every receiver batch, synchronizes CUDA, and +checks optional delta samples. Only after both barriers succeed does the source +commit exact pending baseline bits in background CPU threads. > **Failure boundary:** source baseline commit is transactional, but receiver > writes are in place and are not rolled back. If a transfer fails after a @@ -289,8 +286,8 @@ policy: backend: vllm refit_transport: vllm_s3_sparse # or vllm_zmq_sparse delta_compression: - encoding: overwrite # xor requires bitwise-compatible native loaders - sparse_bucket_size_bytes: 1073741824 + encoding: xor # overwrite is selected automatically for opaque loaders + sparse_bucket_size_bytes: 536870912 colocated: enabled: false vllm_cfg: @@ -310,22 +307,24 @@ same nonempty token on producers and receivers. | Control | Default | |---|---:| -| `NRL_REFIT_S3_EXPORT_CHUNK_BYTES` | 256 MiB | +| `delta_compression.sparse_bucket_size_bytes` | 512 MiB | +| `NRL_REFIT_S3_EXPORT_CHUNK_BYTES` | 64 MiB | | `NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES` | 256 MiB | | `NRL_REFIT_{S3,ZMQ}_ENCODE_WORKERS` | 2-8 from CPU count | | `NRL_REFIT_S3_UPLOAD_WORKERS` | 4-32 from CPU count | | `NRL_REFIT_ZMQ_SEND_WORKERS` | 4 | | `NRL_REFIT_ZMQ_RELAY_PAYLOAD_WORKERS` | 16 | -| `NRL_REFIT_ZMQ_RELAY_FANOUT_WORKERS` | 8-32 from replica count | +| `NRL_REFIT_ZMQ_RELAY_FORWARD_WORKERS` | 8 | +| `NRL_REFIT_ZMQ_RELAY_FANOUT_WORKERS` | receiver endpoints x relay payload workers | | `NRL_REFIT_APPLY_QUEUE_DEPTH` / `NRL_REFIT_APPLY_BATCH_SIZE` | 32 / 8 | | `NRL_REFIT_PARTITION_WORKERS` | 2-8 from CPU count | | `NRL_REFIT_{S3,ZMQ}_ZSTD_THREADS` | 0 | | `NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD` | 0 | Export chunks are capped by `sparse_bucket_size_bytes` and the packed tensor -limit, but intentionally remain smaller than the recommended S3 wire bucket. -Increase one concurrency control at a time; excess parallelism moves the -bottleneck into host memory, Bridge export, relay fanout, or receiver apply. +limit. The S3 defaults were selected by balanced 120B sweeps. Increase one +concurrency control at a time; excess parallelism moves the bottleneck into +host memory, Bridge export, relay-tree forwarding, or receiver apply. ## Metrics and profiling @@ -337,6 +336,7 @@ bottleneck into host memory, Bridge export, relay fanout, or receiver apply. | `REFIT_{S3,ZMQ}_DELTA_CHANGE` | Global changed and total element counts | | `REFIT_RECEIVER_TIMING` | Receiver staging span/wait, batches, apply, and verification | | `REFIT_{S3,ZMQ}_DELTA_VERIFY` | Sampled transmitted-delta accuracy | +| `REFIT_ZMQ_RELAY_FLUSH` | Relay-flush wall time and aggregate fanout service time | | `REFIT_{S3,ZMQ}_GLOBAL_COMMIT` | Successful global flush | `total_s` is producer wall time. Stage fields such as `encode_s`, `s3_put_s`, @@ -344,13 +344,13 @@ and `zmq_send_s` are sums across concurrent tasks and can exceed `total_s`; do not add them as serial phases. Receiver responses also expose node decode/staging, worker deserialization, and native-loader apply time. These are concurrent sums as well, so compare them with receiver wall time rather than -adding them. +adding them. `refit/transfer/relay_flush_s` is the coordinator's ZeroMQ flush +wall time; `fanout_service_s` is an aggregate concurrent service sum. -The `partition` field is `none` for uniquely owned policy-local shards, `names` -for stable name-sharded residual exports, and `chunks` for the full-HF -compatibility path. Synchronizer metrics appear under `refit/delta/*`, -`refit/delta_verify/*`, and `refit/transfer/*` in W&B and other configured -loggers. End-to-end latency is +Chunk counts include every canonical export chunk; changed and total elements +include only the chunks assigned to that producer. Synchronizer metrics appear +under `refit/delta/*`, `refit/delta_verify/*`, and `refit/transfer/*` in W&B and +other configured loggers. End-to-end latency is `timing/train/prepare_for_generation/transfer_and_update_weights`. Nsight ranges cover baseline creation, policy streaming, and vLLM sparse apply. @@ -360,41 +360,38 @@ Relevant thread names start with `nrl-refit-`, `nrl-zmq-`, or ## Development and validation Keep transport changes behind the shared `stream_sparse_delta_payloads()` -pipeline. A transport provides payload delivery and timing only; it must not -duplicate the baseline tracker, codec, receiver queue, or apply logic. Retries -must preserve payload identity and bytes, fan out to every required replica, -and require a successful global flush before source commit. Never retry XOR -after an uncertain or partial receiver apply. +pipeline. A `SparseRefitTransport` provides `name`, `transfer_workers`, +`send(body, payload_id, verification_candidates)`, and worker-local `cleanup()` +only; it must not +duplicate export, baseline tracking, encoding, backpressure, receiver queue, or +apply logic. Retries must preserve payload identity and bytes, fan out to every +required replica, and require a successful global flush before source commit. +Never retry XOR after an uncertain or partial receiver apply. Do not add model-specific placement math or a persistent placement cache. New layouts must work through their native vLLM weight loader and the generic -storage-scoped operation context. Tests should cover packed QKV/MLP columns, -local and remote experts, segmented Mamba views, native transforms, dtype casts, -FP8 bit overwrite, contiguous ranges, and explicit locations. Unknown names, -transformed XOR, and overlapping XOR copies must fail closed. +storage-scoped operation context. Tests should cover split, merged, transposed, +overlapping, transformed, skipped, and dtype-cast copy behavior without naming +model families. FP8 bit overwrite, contiguous ranges, and explicit locations +also require coverage. Unknown names, transformed XOR, and overlapping XOR +copies must fail closed. Codec changes must update encoder and decoder together, preserve 64-bit-safe locations, and commit exact source bits only after global success. Receiver changes must preserve FIFO application, bounded memory, error propagation, flush, CUDA synchronization, and clean shutdown. -Every conversion task must remain represented by a unique local baseline unless -the complete policy-local path is disabled for FP8 parameters, quantization, or -custom FSDP. Those configurations retain the canonical full-HF baseline path. -Direct payload mappings must be rectangular affine shards of the exact HF -tensor. Tests must cover column and row offsets, gated splits, replicated -ownership, nonzero TP/ETP ranks, EP-global expert naming, DP/CP and expert-DP -ownership, transactional baseline updates, global changed-task agreement, -stable residual ownership under filtering, grouped-task expansion, and fallback -for transformed mappings. Do not infer an unknown mapping from its suffix, drop -Bridge task dependencies, or modify Megatron Bridge for a transport-specific -hook. +Do not reintroduce an MCore-local baseline, Bridge mapping duplication, or +model-family projection formulas. Tests must cover complete canonical export, +deterministic chunk ownership, transactional baseline updates, transport +cleanup on success and failure, and unchanged producer/receiver overlap. Do not +modify Megatron Bridge for a transport-specific hook. Run the focused suite: ```bash uv run --extra vllm pytest -q \ - tests/unit/utils/test_weight_transfer_remote_sparse.py \ + tests/unit/utils/test_weight_transfer_stream.py \ tests/unit/models/policy/test_megatron_remote_sparse_refit.py \ tests/unit/models/generation/test_vllm_sparse_refit.py \ tests/unit/weight_sync/test_vllm_remote_sparse_weight_synchronizer.py @@ -403,7 +400,7 @@ uv run --extra vllm pytest -q -m vllm \ tests/unit/models/generation/test_vllm_sparse_delta.py uv run ruff check \ - nemo_rl/utils/weight_transfer_{remote_sparse,sparse_codec,zmq}.py \ + nemo_rl/utils/weight_transfer_{http,sparse_codec,stream,zmq}.py \ nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py \ nemo_rl/models/generation/vllm/vllm_{sparse_refit,sparse_delta}.py \ nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py \ @@ -420,15 +417,29 @@ reload the receiver before retrying. ## Refit bandwidth calculator [`refit_bandwidth_calculator.py`](../../tools/refit_bandwidth_calculator.py) is a -calibrated comparison of the checked-in zstd S3 and ZeroMQ measurements against -a measured H100 NCCL envelope. It is not a general fabric or topology simulator. - -The sparse side evaluates `_SPARSE_LATENCY_FITS` for the requested model size, -transport, and any positive `--changed-pct`. The 3% and 5% zstd fits are -measured calibration points; other densities are explicit extrapolations. The -coefficients model end-to-end latency, not transport bandwidth, so -`--candidate-ethernet-gbps` does not rescale S3 or ZeroMQ. Wire bytes use the -corresponding measured zstd size fit. +projection of the latest zstd S3 and ZeroMQ measurements against a measured +H100 NCCL envelope. It is not a general fabric or topology simulator. + +The sparse side uses the current 247.2 GB canonical-HF XOR results with a +512 MiB sparse bucket. Both S3 and ZeroMQ have matched measured 3% and 5% +anchors. End-to-end time and wire bytes scale linearly by indexed model bytes. +Other model sizes and densities outside the measured points are explicit +projections. Because the sparse values are end-to-end measurements rather than +a link model, `--candidate-ethernet-gbps` does not rescale S3 or ZeroMQ. Text +output reports the 512 MiB calibration, and JSON includes +`sparse_bucket_size_bytes: 536870912`. + +| Transport | 3% | 5% | +|---|---:|---:| +| S3 | 20.234 s measured | 25.790 s measured | +| ZeroMQ tree fanout | 24.095 s measured | 33.776 s measured | + +The measured anchors already include full Bridge export on every participating +rank, deterministic chunk ownership, one aggregate checkpoint-sized canonical +comparison across producers, transport work, and all completion barriers. +Do not multiply comparison bytes by the producer count. The ZeroMQ anchor also +includes its staged-ACK tree flush; the calculator does not model relay +forwarding as an independent bandwidth term. The NCCL side interpolates `_NCCL_ANCHORS` in log model-size space. Those anchors were measured at 400 Gbps per rank and are projected onto the requested @@ -456,10 +467,10 @@ crossover NCCL wins; between them the measured range has no single winner. `--json` emits the same fields for scripts. The production transport applies zstd level 1 to every payload, so the -calculator intentionally has no synthetic raw-compression arm. Treat values -outside the 63.2-1121 GB model range or 3%-5% density range, and any different -topology or parallel mapping, as experiment inputs rather than performance -claims. +calculator intentionally has no synthetic raw-compression arm. Treat every +model size other than 247.2 GB, density outside 3%-5% for either transport, and +any different topology or parallel mapping as a projection rather than a +performance claim. ## Failure guide diff --git a/examples/configs/grpo_math_1B.yaml b/examples/configs/grpo_math_1B.yaml index 7eb30cf5c89..b901812bd06 100644 --- a/examples/configs/grpo_math_1B.yaml +++ b/examples/configs/grpo_math_1B.yaml @@ -340,7 +340,7 @@ policy: stop_token_ids: null stop_strings: null refit_transport: null # Set to "vllm_s3_sparse" or "vllm_zmq_sparse" for remote sparse-delta refit. - delta_compression: null # Remote sparse-delta config; null uses the existing refit path. + delta_compression: null # Set {} for XOR and the 512 MiB bucket; null uses the existing refit path. mcore_generation_config: async_engine: false max_model_len: ${policy.max_total_sequence_length} # Engine-side max sequence length. diff --git a/nemo_rl/models/generation/vllm/config.py b/nemo_rl/models/generation/vllm/config.py index 23e6885b457..56ac513176b 100644 --- a/nemo_rl/models/generation/vllm/config.py +++ b/nemo_rl/models/generation/vllm/config.py @@ -14,6 +14,8 @@ from typing import Any, Literal, NotRequired, TypedDict +from pydantic import BaseModel, PositiveInt + from nemo_rl.models.generation.interfaces import GenerationConfig @@ -58,9 +60,9 @@ class VllmSpecificArgs(TypedDict): reasoning_parser_plugin: NotRequired[str] -class VllmDeltaCompressionConfig(TypedDict): - encoding: Literal["xor", "overwrite"] - sparse_bucket_size_bytes: int +class VllmDeltaCompressionConfig(BaseModel, extra="allow"): + encoding: Literal["xor", "overwrite"] = "xor" + sparse_bucket_size_bytes: PositiveInt = 512 * 1024**2 class VllmConfig(GenerationConfig): diff --git a/nemo_rl/models/generation/vllm/vllm_backend.py b/nemo_rl/models/generation/vllm/vllm_backend.py index 2dc2817c867..8344b77330c 100644 --- a/nemo_rl/models/generation/vllm/vllm_backend.py +++ b/nemo_rl/models/generation/vllm/vllm_backend.py @@ -15,7 +15,7 @@ import re import socket import traceback -from typing import Any, cast +from typing import Any import torch import zmq @@ -175,11 +175,10 @@ def prepare_refit_info(self, state_dict_info: dict[str, Any]) -> None: def prepare_sparse_delta_refit_info( self, state_dict_info: dict[str, tuple[tuple[int, ...], torch.dtype]] - ) -> None: - """Reserve the reusable sparse-refit scratch buffer.""" + ) -> list[str]: + """Reserve scratch space and report weights that require overwrite.""" applier = self._get_sparse_delta_applier() - applier.prewarm(state_dict_info) - applier.discover_native_skips(state_dict_info) + return sorted(applier.discover_native_skips(state_dict_info)) def _maybe_process_fp8_kv_cache(self) -> None: """Process weights after loading for FP8 KV cache (static scales).""" @@ -310,12 +309,10 @@ def load_mtp_weights_from_disk(self, model_path: str) -> bool: return False predictor = draft_model.model - mtp_start_layer_idx = cast(int, predictor.mtp_start_layer_idx) - num_mtp_layers = cast(int, predictor.num_mtp_layers) mtp_layer_indices = set( range( - mtp_start_layer_idx, - mtp_start_layer_idx + num_mtp_layers, + predictor.mtp_start_layer_idx, + predictor.mtp_start_layer_idx + predictor.num_mtp_layers, ) ) weights = _read_mtp_layer_weights_from_checkpoint(model_path, mtp_layer_indices) diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py index e397c154365..d607e8a190e 100644 --- a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py +++ b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py @@ -28,6 +28,7 @@ _TensorViewKey = tuple[int, int, tuple[int, ...], tuple[int, ...]] _LoaderWeight = tuple[str, torch.Tensor, sparse_codec.SparseOperation, int, int | None] +_LoaderObservation = tuple[str, int, bool] def _storage_key(tensor: torch.Tensor) -> int: @@ -64,6 +65,7 @@ def __init__( ] = {} self._xor_spans: dict[int, list[tuple[int, int]]] = {} self.copies = 0 + self.xor_compatible = True def start( self, @@ -80,6 +82,34 @@ def start( self._verification_masks.clear() self._xor_spans.clear() self.copies = 0 + self.xor_compatible = True + + def _observe_xor_copy( + self, destination: torch.Tensor, source: torch.Tensor + ) -> bool: + if ( + _storage_key(source) != self._source_storage + or source.dtype != destination.dtype + ): + self.xor_compatible = False + return False + origin = int(destination.storage_offset()) + extents = [ + (int(size) - 1) * int(stride) + for size, stride in zip( + destination.shape, destination.stride(), strict=True + ) + ] + span = ( + origin + sum(min(0, extent) for extent in extents), + origin + sum(max(0, extent) for extent in extents), + ) + spans = self._xor_spans.setdefault(_storage_key(destination), []) + if any(span[0] <= other[1] and other[0] <= span[1] for other in spans): + self.xor_compatible = False + return False + spans.append(span) + return True def _remember_changed( self, destination: torch.Tensor, changed: torch.Tensor @@ -107,6 +137,7 @@ def __torch_dispatch__( return func(*args, **(kwargs or {})) self.copies += 1 + xor_compatible = self._observe_xor_copy(destination, source) if self._operation == "overwrite": if not source.dtype.is_floating_point: raise RuntimeError("Sparse overwrite requires a floating-point loader.") @@ -138,28 +169,11 @@ def __torch_dispatch__( self._remember_changed(destination, changed) return destination - if _storage_key(source) != self._source_storage: + if not xor_compatible: raise RuntimeError( - "XOR cannot pass through a native loader that transforms its input." + "XOR cannot pass through this native loader without changing semantics." ) - if source.dtype != destination.dtype: - raise RuntimeError("XOR source and target dtypes must match.") source = source.expand_as(destination) - origin = int(destination.storage_offset()) - extents = [ - (int(size) - 1) * int(stride) - for size, stride in zip( - destination.shape, destination.stride(), strict=True - ) - ] - span = ( - origin + sum(min(0, extent) for extent in extents), - origin + sum(max(0, extent) for extent in extents), - ) - spans = self._xor_spans.setdefault(_storage_key(destination), []) - if any(span[0] <= other[1] and other[0] <= span[1] for other in spans): - raise RuntimeError("XOR native loader produced overlapping target copies.") - spans.append(span) destination_bits = sparse_codec.integer_view(destination) source_bits = sparse_codec.integer_view(source) changed = source_bits.ne(0) @@ -202,10 +216,10 @@ def __init__(self, model_runner: Any, device: torch.device) -> None: self._skipped_names: set[str] = set() self._verification: list[tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = [] - def prewarm( + def discover_native_skips( self, state_dict_info: Mapping[str, tuple[tuple[int, ...], torch.dtype]] - ) -> None: - """Reserve a reusable buffer for the largest canonical source tensor.""" + ) -> set[str]: + """Reserve scratch and classify rank-local skips and overwrite weights.""" required = max( (prod(shape) * dtype.itemsize for shape, dtype in state_dict_info.values()), default=0, @@ -214,18 +228,13 @@ def prewarm( self._scratch = torch.empty( required, dtype=torch.uint8, device=self._scratch.device ) - - def discover_native_skips( - self, state_dict_info: Mapping[str, tuple[tuple[int, ...], torch.dtype]] - ) -> None: - """Cache weights that the native loader explicitly skips on this rank.""" pending = [ (name, shape, dtype) for name, (shape, dtype) in state_dict_info.items() if dtype.is_floating_point ] if not pending: - return + return set() def weights() -> Iterator[_LoaderWeight]: for name, shape, dtype in pending: @@ -248,8 +257,13 @@ def weights() -> Iterator[_LoaderWeight]: self._validate_loader_report(loaded, observations, allow_unknown_skips=True) if loaded is not None: self._skipped_names.update( - name for name, copies in observations if copies == 0 + name for name, copies, _ in observations if copies == 0 ) + return { + name + for name, copies, xor_compatible in observations + if copies and not xor_compatible + } def _source_tensor(self, item: dict[str, Any]) -> torch.Tensor: shape = tuple(int(dim) for dim in item["shape"]) @@ -327,24 +341,28 @@ def _load_weights( self, weights: Iterable[_LoaderWeight], verification: list[tuple[torch.Tensor, torch.Tensor, torch.Tensor]], - ) -> tuple[Any, list[tuple[str, int]]]: + ) -> tuple[Any, list[_LoaderObservation]]: mode = _SparseWeightLoadMode(self._target_storages, verification) yielded_names: list[str] = [] - observations: list[tuple[str, int]] = [] + observations: list[_LoaderObservation] = [] def observed_weights() -> Iterator[tuple[str, torch.Tensor]]: active = False for name, source, operation, sample_limit, exact_sentinel in weights: if active: mode.finish() - observations.append((yielded_names[-1], mode.copies)) + observations.append( + (yielded_names[-1], mode.copies, mode.xor_compatible) + ) mode.start(source, operation, sample_limit, exact_sentinel) active = True yielded_names.append(name) yield name, source if active: mode.finish() - observations.append((yielded_names[-1], mode.copies)) + observations.append( + (yielded_names[-1], mode.copies, mode.xor_compatible) + ) with torch.no_grad(), mode: loaded = self.model_runner.model.load_weights(observed_weights()) @@ -355,11 +373,11 @@ def observed_weights() -> Iterator[tuple[str, torch.Tensor]]: @staticmethod def _validate_loader_report( loaded: Any, - observations: list[tuple[str, int]], + observations: list[_LoaderObservation], *, allow_unknown_skips: bool, ) -> None: - copied = sum(copies > 0 for _, copies in observations) + copied = sum(copies > 0 for _, copies, _ in observations) if loaded is None: if not allow_unknown_skips and copied != len(observations): raise RuntimeError( @@ -378,37 +396,60 @@ def _apply_decoded_items( ) -> None: def weights() -> Iterator[_LoaderWeight]: for item, locations, values in items: - if str(item["name"]) in self._skipped_names: - continue yield self._prepare_loader_weight(item, locations, values) loaded, observations = self._load_weights(weights(), self._verification) self._validate_loader_report(loaded, observations, allow_unknown_skips=False) + def _iter_sparse_payload( + self, payload: sparse_codec.TensorPayload + ) -> Iterator[sparse_codec.SparseItem]: + packed_locations, value_groups, metadata = payload + for item in metadata: + if str(item["name"]) in self._skipped_names: + continue + value_start = int(item["value_start"]) + value_end = int(item["value_end"]) + location_dtype = ( + torch.int32 + if prod(item["shape"]) <= torch.iinfo(torch.int32).max + else torch.int64 + ) + yield ( + item, + sparse_codec.sparse_locations_for_item( + item, + packed_locations, + device="cpu", + dtype=location_dtype, + ), + value_groups[int(item["value_group"])][value_start:value_end], + ) + @wrap_with_nvtx_name( "vllm_internal_worker_extension/update_weights_from_decoded_sparse_payload" ) def update_weights_from_decoded_sparse_payload( self, *payloads: bytes | str ) -> dict[str, Any]: - return self._load_decoded_sparse_payloads( + return self._load_sparse_payloads( tuple( io.BytesIO(payload) if isinstance(payload, bytes) else payload for payload in payloads ) ) - def _load_decoded_sparse_payloads( + def _load_sparse_payloads( self, sources: tuple[str | io.BytesIO, ...] ) -> dict[str, Any]: started = time.perf_counter() deserialize_s = [0.0] - def decoded_items() -> Iterator[sparse_codec.DecodedSparseItem]: + def sparse_items() -> Iterator[sparse_codec.SparseItem]: for source in sources: item_started = time.perf_counter() payload = cast( - sparse_codec.DecodedSparsePayload, + sparse_codec.TensorPayload, torch.load( source, map_location="cpu", @@ -417,10 +458,10 @@ def decoded_items() -> Iterator[sparse_codec.DecodedSparseItem]: ), ) deserialize_s[0] += time.perf_counter() - item_started - yield from sparse_codec.iter_decoded_sparse_payload(payload) + yield from self._iter_sparse_payload(payload) item_started = time.perf_counter() - self._apply_decoded_items(decoded_items()) + self._apply_decoded_items(sparse_items()) sparse_apply_s = time.perf_counter() - item_started return { "ok": True, diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py index faacb4d1d1e..f2c23c31a27 100644 --- a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py +++ b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py @@ -15,7 +15,6 @@ """Remote sparse-refit receiver lifecycle for vLLM generation workers.""" import asyncio -import io import os import tempfile import threading @@ -23,7 +22,6 @@ from concurrent.futures import Future, ThreadPoolExecutor from typing import Any, Literal, NamedTuple, cast -import torch import uvicorn from fastapi import FastAPI, Request from fastapi.responses import JSONResponse @@ -35,47 +33,36 @@ _get_node_ip_local, ) from nemo_rl.utils import weight_transfer_sparse_codec as sparse_codec -from nemo_rl.utils.weight_transfer_remote_sparse import ( +from nemo_rl.utils.weight_transfer_http import ( G_VLLM_REFIT_API_KEY_HEADER, G_VLLM_REFIT_FLUSH_PATH, G_VLLM_REFIT_PREPARE_PATH, G_VLLM_REFIT_S3_MANIFEST_PATH, + G_VLLM_REFIT_ZMQ_FLUSH_PATH, + merge_vllm_refit_metrics, + vllm_refit_api_key, +) +from nemo_rl.utils.weight_transfer_stream import ( decode_sparse_payload, download_s3_refit_payload, - merge_vllm_refit_metrics, refit_env_int, - vllm_refit_api_key, ) from nemo_rl.utils.weight_transfer_zmq import ( G_VLLM_REFIT_CHECKSUM_HEADER, G_VLLM_REFIT_PAYLOAD_HEADER, G_VLLM_REFIT_PRODUCER_HEADER, G_VLLM_REFIT_TRANSFER_HEADER, + G_VLLM_REFIT_VERIFICATION_HEADER, G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, ZmqSparseRefitServer, ) -def _decode_staged_payload( - serialized: bytes, -) -> tuple[sparse_codec.DecodedSparsePayload, int]: - payload = cast( - sparse_codec.TensorPayload, - torch.load(io.BytesIO(serialized), map_location="cpu", weights_only=True), - ) - return ( - sparse_codec.decode_sparse_tensor_payload_for_staging(payload), - sum(int(item.get("verification_samples", 0)) for item in payload[2]), - ) - - class _StagedSparsePayload(NamedTuple): path: str started_at: float finished_at: float - deserialize_s: float save_s: float - candidates: int def _stage_sparse_payload( @@ -83,15 +70,14 @@ def _stage_sparse_payload( staging_dir: str, ) -> _StagedSparsePayload: started_at = time.perf_counter() - decoded, candidates = _decode_staged_payload(serialized) - deserialize_s = time.perf_counter() - started_at descriptor, path = tempfile.mkstemp( prefix="nemo_rl_refit_", suffix=".pt", dir=staging_dir ) - os.close(descriptor) started = time.perf_counter() try: - torch.save(decoded, path) + with os.fdopen(descriptor, "wb") as handle: + handle.write(serialized) + save_s = time.perf_counter() - started except Exception: os.unlink(path) raise @@ -100,9 +86,7 @@ def _stage_sparse_payload( path, started_at, finished_at, - deserialize_s, - finished_at - started, - candidates, + save_s, ) @@ -171,6 +155,7 @@ def _enqueue_sparse_payload_apply( payload: bytes, payload_key: tuple[str, int, int], checksum: str, + verification_candidates: int = 0, ) -> dict[str, Any]: completed: list[Future[dict[str, Any]]] = [] with self._refit_apply_queue_condition: @@ -190,6 +175,7 @@ def _enqueue_sparse_payload_apply( completed.append(self._refit_apply_futures.pop(0)) response = self._collect_refit_apply_results(completed) self._refit_seen_payloads[payload_key] = checksum + self._refit_verification_candidates += verification_candidates pending: bytes | Future[_StagedSparsePayload] = payload if self._refit_workers_share_node: pending = self._refit_partition_executor.submit( @@ -261,23 +247,10 @@ def update_weights_from_serialized_sparse_payloads( serialized_payloads: tuple[bytes, ...], ) -> dict[str, Any]: """Apply a FIFO batch of sparse deltas through one collective RPC.""" - - def decode_for_rpc(serialized: bytes) -> tuple[bytes, int]: - decoded, candidates = _decode_staged_payload(serialized) - buffer = io.BytesIO() - torch.save(decoded, buffer) - return buffer.getvalue(), candidates - - decoded_payloads = list( - self._refit_partition_executor.map(decode_for_rpc, serialized_payloads) - ) - self._refit_verification_candidates += sum( - candidates for _, candidates in decoded_payloads - ) response = self._refit_collective_response( self._refit_collective_rpc( "update_weights_from_decoded_sparse_payload", - tuple(payload for payload, _ in decoded_payloads), + serialized_payloads, ) ) response["payloads"] = len(serialized_payloads) @@ -299,9 +272,6 @@ def update_weights_from_staged_sparse_payloads( if stage_error is not None: raise stage_error stage_wait_s = time.perf_counter() - started - self._refit_verification_candidates += sum( - payload.candidates for payload in staged - ) try: response = self._refit_collective_response( self._refit_collective_rpc( @@ -318,9 +288,7 @@ def update_weights_from_staged_sparse_payloads( os.unlink(payload.path) worker_total_s = float(response.get("receiver_total_s", 0.0)) response.update( - receiver_node_deserialize_s=max( - (payload.deserialize_s for payload in staged), default=0.0 - ), + receiver_node_deserialize_s=0.0, receiver_stage_s=( max(payload.finished_at for payload in staged) - min(payload.started_at for payload in staged) @@ -390,16 +358,22 @@ def _prepare_sparse_refit_info(self, request: dict[str, Any]) -> dict[str, Any]: name: (tuple(shape), sparse_codec.dtype_from_name(dtype)) for name, (shape, dtype) in request["tensors"].items() } - self._refit_collective_rpc( + worker_results = self._refit_collective_rpc( "prepare_sparse_delta_refit_info", (state_dict_info,) ) + overwrite_names = sorted({name for names in worker_results for name in names}) seconds = time.perf_counter() - started print( f"REFIT_RECEIVER_PREWARM tensors={len(state_dict_info)} " - f"seconds={seconds:.3f}", + f"overwrite_tensors={len(overwrite_names)} seconds={seconds:.3f}", flush=True, ) - return {"ok": True, "tensors": len(state_dict_info), "seconds": seconds} + return { + "ok": True, + "tensors": len(state_dict_info), + "overwrite_names": overwrite_names, + "seconds": seconds, + } async def _apply_s3_manifest_payload( self, @@ -415,6 +389,7 @@ async def _apply_s3_manifest_payload( body, (key, -1, -1), checksum, + int(manifest["verification_candidates"]), ) result["receiver_s3_download_s"] = download_s return result @@ -425,7 +400,16 @@ async def _apply_zmq_payload(self, raw_request: Any) -> dict[str, Any]: producer_id = int(headers.get(G_VLLM_REFIT_PRODUCER_HEADER, "-1")) payload_id = int(headers.get(G_VLLM_REFIT_PAYLOAD_HEADER, "-1")) checksum = headers.get(G_VLLM_REFIT_CHECKSUM_HEADER, "") - if not transfer_id or producer_id < 0 or payload_id < 0 or not checksum: + verification_candidates = int( + headers.get(G_VLLM_REFIT_VERIFICATION_HEADER, "-1") + ) + if ( + not transfer_id + or producer_id < 0 + or payload_id < 0 + or not checksum + or verification_candidates < 0 + ): raise ValueError("Missing or invalid ZeroMQ sparse refit payload headers.") compressed = await raw_request.body() started = time.perf_counter() @@ -440,6 +424,7 @@ async def _apply_zmq_payload(self, raw_request: Any) -> dict[str, Any]: payload, (transfer_id, producer_id, payload_id), checksum, + verification_candidates, ) result["receiver_zmq_decode_s"] = decode_s return result @@ -450,7 +435,7 @@ def setup_api_server(self, app: Any) -> None: async def respond( raw_request: Request, - action: Literal["prepare", "s3", "flush", "zmq"], + action: Literal["prepare", "s3", "flush", "zmq", "zmq_flush"], ) -> JSONResponse: if cfg["vllm_cfg"]["async_engine"]: self._refit_async_loop = asyncio.get_running_loop() @@ -473,6 +458,13 @@ async def respond( ) elif action == "zmq": result = await self._apply_zmq_payload(raw_request) + elif action == "zmq_flush": + body = await raw_request.json() + result = await asyncio.to_thread( + self.flush_zmq_sparse_refit_relay, + str(body["transfer_id"]), + int(body.get("expected_payloads", 0)), + ) else: result = await asyncio.to_thread(self._flush_queued_sparse_payloads) except Exception as exc: @@ -482,7 +474,9 @@ async def respond( status_code=200 if result.get("ok") is True else 500, ) - def endpoint(action: Literal["prepare", "s3", "flush", "zmq"]): + def endpoint( + action: Literal["prepare", "s3", "flush", "zmq", "zmq_flush"], + ): async def handle(raw_request: Request) -> JSONResponse: return await respond(raw_request, action) @@ -493,6 +487,7 @@ async def handle(raw_request: Request) -> JSONResponse: (G_VLLM_REFIT_PREPARE_PATH, "prepare"), (G_VLLM_REFIT_FLUSH_PATH, "flush"), (G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, "zmq"), + (G_VLLM_REFIT_ZMQ_FLUSH_PATH, "zmq_flush"), ): app.add_api_route(path, endpoint(action), methods=["POST"]) @@ -522,11 +517,31 @@ def start_zmq_sparse_refit_relay(self, refit_urls: list[str]) -> str: print(f"Starting vLLM ZeroMQ refit relay on {address}", flush=True) return address + def configure_zmq_sparse_refit_relay(self, relay_addresses: list[str]) -> None: + if self._zmq_refit_server is None: + raise RuntimeError("ZeroMQ sparse refit relay is not running.") + local_refit_url = self.report_refit_server_base_url() + if local_refit_url is None: + raise RuntimeError("Local vLLM sparse refit endpoint is unavailable.") + server, own_address = self._zmq_refit_server + server.configure_tree( + relay_addresses, + own_address=own_address, + local_refit_url=local_refit_url, + ) + def stop_zmq_sparse_refit_relay(self) -> None: if self._zmq_refit_server is not None: self._zmq_refit_server[0].close() self._zmq_refit_server = None + def flush_zmq_sparse_refit_relay( + self, transfer_id: str, expected_payloads: int = 0 + ) -> dict[str, Any]: + if self._zmq_refit_server is None: + raise RuntimeError("ZeroMQ sparse refit relay is not running.") + return self._zmq_refit_server[0].flush(transfer_id, expected_payloads) + def _setup_vllm_refit_server(self) -> None: app = FastAPI() self.setup_api_server(app) diff --git a/nemo_rl/models/generation/vllm/vllm_worker.py b/nemo_rl/models/generation/vllm/vllm_worker.py index e3e1791272b..316f8c6bbf3 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker.py +++ b/nemo_rl/models/generation/vllm/vllm_worker.py @@ -668,6 +668,12 @@ def start_zmq_sparse_refit_relay(self, refit_urls: list[str]) -> str: raise RuntimeError("Remote sparse refit is not enabled for this worker.") return receiver.start_zmq_sparse_refit_relay(refit_urls) + def configure_zmq_sparse_refit_relay(self, relay_addresses: list[str]) -> None: + receiver = self._sparse_refit_receiver + if receiver is None: + raise RuntimeError("Remote sparse refit is not enabled for this worker.") + receiver.configure_zmq_sparse_refit_relay(relay_addresses) + def stop_zmq_sparse_refit_relay(self) -> None: receiver = self._sparse_refit_receiver if receiver is not None: diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index a08bc0c6596..9f775e1db32 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -1823,6 +1823,7 @@ def stream_remote_sparse_weights( timeout_s: float, shard_rank: int, shard_count: int, + overwrite_names: list[str], ) -> dict[str, int]: return self._require_remote_sparse_refit().stream( transport, @@ -1832,6 +1833,7 @@ def stream_remote_sparse_weights( timeout_s=timeout_s, shard_rank=shard_rank, shard_count=shard_count, + overwrite_names=overwrite_names, ) def _require_remote_sparse_refit(self) -> Any: @@ -1906,7 +1908,6 @@ def calculate_size_in_bytes(param, tp_size, ep_size): def _iter_params_with_optional_kv_scales( self, kv_scales: Optional[dict[str, float]] = None, - conversion_tasks: Optional[list[Any]] = None, ) -> Iterator[tuple[str, torch.Tensor]]: """Yield exported HF parameters and optionally append FP8 KV/Q scale tensors. @@ -1917,12 +1918,10 @@ def _iter_params_with_optional_kv_scales( get_vllm_qkv_scale_names, ) - if conversion_tasks is None: - conversion_tasks = self.refit_conversion_tasks base_iter = self.megatron_bridge.export_hf_weights( [self.model], show_progress=False, - conversion_tasks=conversion_tasks, + conversion_tasks=self.refit_conversion_tasks, ) # Yield the original parameters first. diff --git a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py index 762debd8538..3da8550596f 100644 --- a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py +++ b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py @@ -12,371 +12,29 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Optional remote sparse-refit state owned by a Megatron policy worker.""" +"""Canonical Hugging Face sparse-refit state for a Megatron policy worker.""" -import re -from collections.abc import Iterable, Mapping -from concurrent.futures import ThreadPoolExecutor -from functools import cache, partial -from typing import Any, cast +from collections.abc import Iterator +from typing import Any import torch -from nemo_rl.utils.weight_transfer_remote_sparse import ( +from nemo_rl.models.generation.vllm.config import VllmDeltaCompressionConfig +from nemo_rl.utils.weight_transfer_sparse_codec import DeltaCompressionTracker +from nemo_rl.utils.weight_transfer_stream import ( init_sparse_delta_baseline_from_iterator, - sparse_name_shard, stream_sparse_delta_payloads_via_s3_manifest, ) -from nemo_rl.utils.weight_transfer_sparse_codec import ( - DeltaCompressionTracker, - SparseShardProjection, -) from nemo_rl.utils.weight_transfer_zmq import stream_sparse_delta_payloads_via_zmq -_UNSUPPORTED = 0 -_COLUMN = 1 -_ROW = 2 -_REPLICATED = 3 -_GATED = 4 - class MegatronRemoteSparseRefit: - def __init__(self, worker: Any, delta_config: Mapping[str, Any]) -> None: + def __init__(self, worker: Any, delta_config: VllmDeltaCompressionConfig) -> None: self._worker = worker - self._delta_config = delta_config - residual_config = dict(delta_config) - if residual_config["encoding"] == "xor": - residual_config["encoding"] = "overwrite" - # Residual Bridge outputs need a canonical-HF overwrite baseline; the - # optional policy tracker owns MCore-local projections and can retain XOR. - self._tracker = DeltaCompressionTracker(residual_config) - self._policy_tracker: DeltaCompressionTracker | None = None - self._local_tensors: list[tuple[str, torch.Tensor]] = [] - self._misc_local_tensors: list[tuple[str, torch.Tensor]] = [] - self._misc_conversion_tasks: list[Any] | None = None - self._filter_misc_tasks = False - - @staticmethod - @cache - def _bridge_mapping_types() -> tuple[Any, dict[Any, int]]: - # Bridge is optional outside Megatron workers, so keep these imports local. - from megatron.bridge.models.conversion.param_mapping import ( - AutoMapping, - ColumnParallelMapping, - DirectMapping, - GatedMLPMapping, - ReplicatedMapping, - RowParallelMapping, - ) - - return AutoMapping, { - ColumnParallelMapping: _COLUMN, - DirectMapping: _REPLICATED, - GatedMLPMapping: _GATED, - ReplicatedMapping: _REPLICATED, - RowParallelMapping: _ROW, - } - - @staticmethod - def _all_reduce_max(values: list[int]) -> list[int]: - if not values or not torch.distributed.is_initialized(): - return values - backend = str(torch.distributed.get_backend()).lower() - device = ( - torch.device("cuda", torch.cuda.current_device()) - if backend.endswith("nccl") - else torch.device("cpu") - ) - reduced = torch.tensor(values, dtype=torch.int32, device=device) - torch.distributed.all_reduce(reduced, op=torch.distributed.ReduceOp.MAX) - return reduced.cpu().tolist() - - @classmethod - def _local_mapping_kind(cls, task: Any) -> int: - AutoMapping, mapping_kinds = cls._bridge_mapping_types() - mapping = task.mapping - if kind := mapping_kinds.get(type(mapping)): - return kind - if ( - type(mapping) is AutoMapping - and mapping.permute_dims is None - and task.megatron_module is not None - ): - return { - "column": _COLUMN, - "row": _ROW, - "replicated": _REPLICATED, - }.get(mapping._detect_parallelism_type(task.megatron_module), _UNSUPPORTED) - return _UNSUPPORTED - - def _bridge_exports_are_identity(self) -> bool: - bridge = getattr(self._worker, "megatron_bridge", None) - if bridge is None: - return True - from megatron.bridge.models.conversion.model_bridge import MegatronModelBridge - - model_bridge = bridge._model_bridge - return ( - type(model_bridge).maybe_modify_converted_hf_weight - is MegatronModelBridge.maybe_modify_converted_hf_weight - ) + self._tracker = DeltaCompressionTracker(delta_config.model_dump()) - @staticmethod - def _is_padded_or_tied_weight(task: Any) -> bool: - hf_names = ( - task.mapping.hf_param.values() - if isinstance(task.mapping.hf_param, dict) - else (task.mapping.hf_param,) - ) - return task.global_param_name.endswith( - ("embedding.word_embeddings.weight", "output_layer.weight") - ) or any( - str(name).endswith( - ("embed_tokens.weight", "embeddings.weight", "lm_head.weight") - ) - for name in hf_names - ) - - @staticmethod - def _owns_policy_local_task(task: Any, *, replicated: bool = False) -> bool: - if not torch.distributed.is_initialized(): - return True - from megatron.core import parallel_state - - if task.mapping.is_expert: - # Expert-DP includes DP replicas and TP replicas when ETP < TP. - replica_rank = parallel_state.get_expert_data_parallel_rank() - replica_count = parallel_state.get_expert_data_parallel_world_size() - else: - replica_rank = parallel_state.get_data_parallel_rank( - with_context_parallel=True - ) - replica_count = parallel_state.get_data_parallel_world_size( - with_context_parallel=True - ) - if replicated: - replica_rank = ( - replica_rank * task.mapping.tp_size + task.mapping.tp_rank - ) - replica_count *= task.mapping.tp_size - return replica_rank == sparse_name_shard(task.global_param_name, replica_count) - - @classmethod - def _task_ownership(cls, task: Any, kind: int) -> tuple[torch.Tensor | None, bool]: - tensor = task.param_weight - replicated = kind == _REPLICATED or ( - kind == _ROW and tensor is not None and tensor.ndim == 1 - ) - return ( - tensor - if tensor is not None - and cls._owns_policy_local_task(task, replicated=replicated) - else None, - replicated, - ) - - def _policy_local_path_is_safe(self) -> bool: - config = getattr(self._worker, "cfg", {}) - ddp_config = config.get("megatron_cfg", {}).get( - "distributed_data_parallel_config", {} - ) - fp8_cfg = getattr(self._worker, "fp8_cfg", None) - return ( - config.get("quant_cfg") is None - and not ddp_config.get("use_custom_fsdp", False) - and (not fp8_cfg or not fp8_cfg.get("fp8_param", False)) - ) - - @staticmethod - def _canonical_hf_name(task: Any, name: str) -> str: - mapping = task.mapping - if not mapping.is_expert or mapping.ep_size == 1: - return name - - match = re.search(r"(?<=\.experts\.)\d+(?=\.)", name) - config = getattr(task.megatron_module, "config", None) - num_experts = getattr(config, "num_moe_experts", None) - if match is None or not isinstance(num_experts, int): - raise ValueError(f"Cannot project expert parameter {name!r}.") - if num_experts % mapping.ep_size: - raise ValueError( - f"Expert count {num_experts} is not divisible by EP size " - f"{mapping.ep_size}." - ) - experts_per_rank = num_experts // mapping.ep_size - expert = int(match.group()) % experts_per_rank - expert += experts_per_rank * mapping.ep_rank - return f"{name[: match.start()]}{expert}{name[match.end() :]}" - - @staticmethod - def _projection( - name: str, - tensor: torch.Tensor, - *, - shard_dim: int, - shard_rank: int, - shard_count: int, - ) -> SparseShardProjection: - if tensor.ndim <= shard_dim: - raise ValueError(f"Cannot shard {name!r} on dimension {shard_dim}.") - global_shape = list(tensor.shape) - global_shape[shard_dim] *= shard_count - return SparseShardProjection( - name, - tuple(global_shape), - shard_dim, - tensor.shape[shard_dim] * shard_rank, - ) - - @classmethod - def _task_local_tensors( - cls, task: Any, kind: int, *, identity_export: bool = True - ) -> list[tuple[str, torch.Tensor, SparseShardProjection]] | None: - mapping = task.mapping - hf_param = mapping.hf_param - if ( - kind == _UNSUPPORTED - or not identity_export - or getattr(mapping, "is_grouped_export", False) - or getattr(mapping, "is_adapter", False) - or cls._is_padded_or_tied_weight(task) - ): - return None - if kind == _GATED: - if not isinstance(hf_param, dict) or set(hf_param) != { - "gate", - "up", - }: - return None - elif not isinstance(hf_param, str): - return None - tensor, replicated = cls._task_ownership(task, kind) - if tensor is None: - return [] - if kind == _GATED: - gate, up = torch.chunk(tensor, 2, dim=0) - return [ - ( - f"{task.global_param_name}:{role}", - value, - cls._projection( - cls._canonical_hf_name( - task, str(cast(dict[str, Any], hf_param)[role]) - ), - value, - shard_dim=0, - shard_rank=mapping.tp_rank, - shard_count=mapping.tp_size, - ), - ) - for role, value in (("gate", gate), ("up", up)) - ] - - name = cls._canonical_hf_name(task, cast(str, hf_param)) - projection = ( - SparseShardProjection(name, tuple(tensor.shape)) - if replicated - else cls._projection( - name, - tensor, - shard_dim=0 if kind == _COLUMN else 1, - shard_rank=mapping.tp_rank, - shard_count=mapping.tp_size, - ) - ) - return [(task.global_param_name, tensor, projection)] - - def _prepare_paths(self) -> None: - if self._misc_conversion_tasks is not None: - return - tasks = [ - task - for task in self._worker.megatron_bridge.get_conversion_tasks( - [self._worker.model] - ) - if task is not None - ] - misc_tasks: list[Any] = [] - self._misc_conversion_tasks = misc_tasks - if not tasks or not self._policy_local_path_is_safe(): - misc_tasks.extend(tasks) - return - kinds = self._all_reduce_max([self._local_mapping_kind(task) for task in tasks]) - identity_export = self._bridge_exports_are_identity() - self._filter_misc_tasks = identity_export - projections = {} - for task, kind in zip(tasks, kinds, strict=True): - local_tensors = self._task_local_tensors( - task, kind, identity_export=identity_export - ) - if local_tensors is not None: - for key, tensor, projection in local_tensors: - if key in projections: - raise ValueError( - f"Duplicate policy-local sparse shard {key!r}." - ) - projections[key] = projection - self._local_tensors.append((key, tensor)) - continue - - task_index = len(misc_tasks) - misc_tasks.append(task) - tensor, _ = self._task_ownership(task, kind) - if tensor is not None: - key = f"{task_index}:{task.global_param_name}" - self._misc_local_tensors.append((key, tensor)) - - self._policy_tracker = DeltaCompressionTracker( - self._delta_config, projections=projections - ) - - def _iter_misc_params( - self, conversion_tasks: list[Any] | None = None - ) -> Iterable[tuple[str, torch.Tensor]]: - tasks = ( - self._misc_conversion_tasks - if conversion_tasks is None - else conversion_tasks - ) - return self._worker._iter_params_with_optional_kv_scales(conversion_tasks=tasks) - - def _changed_misc_tasks(self) -> tuple[list[Any], int, int]: - assert self._misc_conversion_tasks is not None - changed_keys: set[str] = set() - changed = total = 0 - if self._misc_local_tensors: - assert self._policy_tracker is not None - changed_keys, changed, total = self._policy_tracker.prepare_change_summary( - self._misc_local_tensors - ) - flags = [0] * len(self._misc_conversion_tasks) - for key in changed_keys: - flags[int(key.partition(":")[0])] = 1 - flags = self._all_reduce_max(flags) - if not any(flags): - tasks = [] - elif not self._filter_misc_tasks: - tasks = self._misc_conversion_tasks - else: - grouped_keys = { - task.mapping.group_key - for task, task_changed in zip( - self._misc_conversion_tasks, flags, strict=True - ) - if task_changed and getattr(task.mapping, "is_grouped_export", False) - } - tasks = [ - task - for task, task_changed in zip( - self._misc_conversion_tasks, flags, strict=True - ) - if task_changed - or ( - getattr(task.mapping, "is_grouped_export", False) - and task.mapping.group_key in grouped_keys - ) - ] - return tasks, changed, total + def _iter_params(self) -> Iterator[tuple[str, torch.Tensor]]: + return self._worker._iter_params_with_optional_kv_scales() def initialize_baseline( self, @@ -385,45 +43,17 @@ def initialize_baseline( shard_count: int, transport: str, ) -> dict[str, tuple[tuple[int, ...], torch.dtype]]: - self._prepare_paths() - snapshot = partial( - init_sparse_delta_baseline_from_iterator, + init_sparse_delta_baseline_from_iterator( + self._iter_params(), + delta_tracker=self._tracker, shard_rank=shard_rank, shard_count=shard_count, transport=transport, ) - policy_tracker = self._policy_tracker - if policy_tracker is None: - snapshot(self._iter_misc_params(), delta_tracker=self._tracker) - return self.refit_info() - - with ThreadPoolExecutor( - max_workers=1, thread_name_prefix="nrl-refit-policy-local" - ) as executor: - local_future = executor.submit( - snapshot, - self._local_tensors + self._misc_local_tensors, - delta_tracker=policy_tracker, - partition="none", - ) - snapshot( - self._iter_misc_params(), - delta_tracker=self._tracker, - partition="names", - ) - local_future.result() - return self.refit_info() - - def refit_info(self) -> dict[str, tuple[tuple[int, ...], torch.dtype]]: - info = { + return { name: (tuple(tensor.shape), tensor.dtype) for name, tensor in self._tracker.baseline.items() } - if self._policy_tracker is not None: - for name, tensor in self._local_tensors: - projection = self._policy_tracker.projections[name] - info[projection.name] = (projection.global_shape, tensor.dtype) - return info def stream( self, @@ -435,66 +65,29 @@ def stream( timeout_s: float, shard_rank: int, shard_count: int, + overwrite_names: list[str], ) -> dict[str, int]: - self._prepare_paths() streamer = { "s3": stream_sparse_delta_payloads_via_s3_manifest, "zmq": stream_sparse_delta_payloads_via_zmq, }[transport] - send = partial( - streamer, + self._tracker.overwrite_names = frozenset(overwrite_names) + result = streamer( + self._iter_params(), + delta_tracker=self._tracker, + transfer_id=transfer_id, refit_targets=targets, api_key_env_var=api_key_env_var, timeout_s=timeout_s, shard_rank=shard_rank, shard_count=shard_count, ) - policy_tracker = self._policy_tracker - if policy_tracker is None: - result = send( - self._iter_misc_params(), - delta_tracker=self._tracker, - transfer_id=transfer_id, - ) - else: - with ThreadPoolExecutor( - max_workers=1, thread_name_prefix="nrl-refit-policy-local" - ) as executor: - local_future = None - if self._local_tensors: - local_future = executor.submit( - send, - self._local_tensors, - delta_tracker=policy_tracker, - transfer_id=f"{transfer_id}-local", - partition="none", - ) - changed_tasks, misc_changed, misc_total = self._changed_misc_tasks() - misc_result = send( - self._iter_misc_params(changed_tasks), - delta_tracker=self._tracker, - transfer_id=f"{transfer_id}-misc", - partition="names", - ) - local_result = ( - local_future.result() - if local_future is not None - else {"payloads": 0, "changed_elements": 0, "total_elements": 0} - ) - result = dict( - payloads=int(local_result["payloads"]) + int(misc_result["payloads"]), - changed_elements=int(local_result["changed_elements"]) + misc_changed, - total_elements=int(local_result["total_elements"]) + misc_total, - ) if torch.cuda.is_available(): torch.cuda.synchronize() return result def finish(self, succeeded: bool) -> None: - for tracker in (self._tracker, self._policy_tracker): - if tracker is None: - continue - if succeeded: - tracker.on_sync_succeeded() - else: - tracker.on_sync_failed() + if succeeded: + self._tracker.on_sync_succeeded() + else: + self._tracker.on_sync_failed() diff --git a/nemo_rl/utils/weight_transfer_http.py b/nemo_rl/utils/weight_transfer_http.py new file mode 100644 index 00000000000..29168b2e8da --- /dev/null +++ b/nemo_rl/utils/weight_transfer_http.py @@ -0,0 +1,141 @@ +# 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. + +"""HTTP control-plane utilities shared by remote sparse-refit transports.""" + +import os +import threading +from collections.abc import Iterable, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from functools import cache +from typing import Any + +import requests +from urllib3.util.retry import Retry + +G_VLLM_REFIT_S3_MANIFEST_PATH = "/nemo-rl/refit/s3-manifest" +G_VLLM_REFIT_PREPARE_PATH = "/nemo-rl/refit/prepare" +G_VLLM_REFIT_FLUSH_PATH = "/nemo-rl/refit/flush" +G_VLLM_REFIT_ZMQ_FLUSH_PATH = "/nemo-rl/refit/zmq-flush" +G_VLLM_REFIT_API_KEY_HEADER = "x-nemo-rl-refit-key" +_HTTP_LOCAL = threading.local() +_HTTP_ADAPTER = requests.adapters.HTTPAdapter( + pool_connections=64, + pool_maxsize=64, + max_retries=Retry( + total=3, + backoff_factor=0.25, + status_forcelist=(502, 503, 504), + allowed_methods={"POST"}, + ), +) + + +def vllm_refit_endpoints(base_urls: Sequence[str], path: str) -> list[str]: + return list( + dict.fromkeys( + f"{url.strip().rstrip('/')}{path}" for url in base_urls if url.strip() + ) + ) + + +def vllm_refit_api_key(api_key_env_var: str | None) -> str | None: + if not api_key_env_var: + return None + token = os.environ.get(api_key_env_var) + if not token: + raise RuntimeError( + "vLLM sparse refit API key env var " + f"{api_key_env_var!r} is configured but unset or empty." + ) + return token + + +def refit_http_session() -> requests.Session: + session = getattr(_HTTP_LOCAL, "session", None) + if session is None: + session = requests.Session() + session.mount("http://", _HTTP_ADAPTER) + session.mount("https://", _HTTP_ADAPTER) + _HTTP_LOCAL.session = session + return session + + +@cache +def _http_executor(workers: int) -> ThreadPoolExecutor: + return ThreadPoolExecutor(max_workers=workers, thread_name_prefix="nrl-refit-http") + + +def post_vllm_refit_endpoints( + endpoint_urls: Sequence[str], + body: Mapping[str, Any] | bytes, + *, + api_key: str | None, + timeout_s: float, + headers: Mapping[str, str] | None = None, + executor: ThreadPoolExecutor | None = None, +) -> list[dict[str, Any]]: + request_headers = dict(headers or {}) + if api_key: + request_headers[G_VLLM_REFIT_API_KEY_HEADER] = api_key + request_kwargs: dict[str, Any] = ( + {"data": body} if isinstance(body, bytes) else {"json": body} + ) + + def post(url: str) -> dict[str, Any]: + response = refit_http_session().post( + url, + **request_kwargs, + headers=request_headers, + timeout=timeout_s, + ) + try: + result: dict[str, Any] = response.json() if response.content else {} + except requests.exceptions.JSONDecodeError: + result = {} + if response.status_code >= 400 or result.get("ok") is not True: + raise RuntimeError( + f"vLLM refit failed for {url}: HTTP {response.status_code}: " + f"{response.text[:512]}" + ) + return result + + pool = executor or _http_executor(len(endpoint_urls)) + return list(pool.map(post, endpoint_urls)) + + +def merge_vllm_refit_metrics( + result: dict[str, Any], + metrics: Iterable[Mapping[str, Any]], + *, + maximum: bool, + candidate_maximum: bool | None = None, +) -> dict[str, Any]: + for metric in metrics: + for key, value in metric.items(): + if key.startswith("receiver_") and key.endswith("_s"): + number, use_maximum = float(value), maximum + elif candidate_maximum is not None and key.startswith("verification_"): + number = value + use_maximum = key == "verification_max_abs" or ( + key == "verification_candidates" and candidate_maximum + ) + else: + continue + if key in result: + number = ( + max(result[key], number) if use_maximum else result[key] + number + ) + result[key] = number + return result diff --git a/nemo_rl/utils/weight_transfer_sparse_codec.py b/nemo_rl/utils/weight_transfer_sparse_codec.py index 72caa89e3f8..c53216c441f 100644 --- a/nemo_rl/utils/weight_transfer_sparse_codec.py +++ b/nemo_rl/utils/weight_transfer_sparse_codec.py @@ -17,8 +17,6 @@ import threading from collections.abc import Iterable, Mapping from concurrent.futures import ThreadPoolExecutor -from dataclasses import dataclass -from math import prod from typing import Any, Literal import numpy as np @@ -30,12 +28,7 @@ SparseInfo = tuple[str, torch.Tensor, torch.Tensor, torch.Tensor, SparseOperation] TensorPayload = tuple[torch.Tensor, tuple[torch.Tensor, ...], list[dict[str, Any]]] PreparedTensorPayload = tuple[TensorPayload, int, int] -DecodedSparseItem = tuple[dict[str, Any], torch.Tensor, torch.Tensor] -DecodedSparsePayload = tuple[ - tuple[torch.Tensor, torch.Tensor], - tuple[torch.Tensor, ...], - list[dict[str, Any]], -] +SparseItem = tuple[dict[str, Any], torch.Tensor, torch.Tensor] _INTEGER_DTYPE_BY_SIZE = { 1: torch.uint8, @@ -86,32 +79,6 @@ def finish(self) -> TensorPayload: return indices, values, self.metadata -@dataclass(frozen=True) -class SparseShardProjection: - """Map one local tensor shard into a canonical HF tensor.""" - - name: str - global_shape: tuple[int, ...] - shard_dim: int | None = None - offset: int = 0 - - def map_locations( - self, locations: torch.Tensor, local_shape: tuple[int, ...] - ) -> torch.Tensor: - if self.shard_dim is None: - return locations - if self.shard_dim == 0: - return locations + self.offset * prod(local_shape[1:]) - inner = prod(local_shape[self.shard_dim + 1 :]) - local_slab = local_shape[self.shard_dim] * inner - outer = torch.div(locations, local_slab, rounding_mode="floor") - return ( - locations - + self.offset * inner - + outer * (prod(self.global_shape[self.shard_dim :]) - local_slab) - ) - - def integer_dtype_for_element_size(element_size: int) -> torch.dtype: try: return _INTEGER_DTYPE_BY_SIZE[element_size] @@ -217,77 +184,16 @@ def sparse_locations_for_item( return torch.arange(start, start + count, dtype=dtype, device=device) index_start, index_end = int(item["index_start"]), int(item["index_end"]) - raw = ( - packed_locations[index_start:index_end] - .detach() - .cpu() - .numpy() - .astype(np.uint8, copy=False) - .tobytes() - ) - delta_dtype = {2: np.uint16, 4: np.uint32, 8: np.uint64}[len(raw) // count] + raw = packed_locations[index_start:index_end].detach().cpu().numpy() + delta_dtype = {2: np.uint16, 4: np.uint32, 8: np.uint64}[raw.size // count] location_dtype = np.int32 if dtype == torch.int32 else np.int64 - deltas = np.frombuffer(raw, dtype=delta_dtype).astype(location_dtype, copy=False) - locations = np.cumsum(deltas + 1, dtype=location_dtype) - 1 + locations = raw.view(delta_dtype).astype(location_dtype, copy=False) + locations += 1 + np.cumsum(locations, out=locations) + locations -= 1 return torch.from_numpy(locations).to(device=device) -def _merge_tensor_parts(parts: list[torch.Tensor], dtype: torch.dtype) -> torch.Tensor: - if not parts: - return torch.empty(0, dtype=dtype) - return parts[0] if len(parts) == 1 else torch.cat(parts) - - -def decode_sparse_tensor_payload_for_staging( - payload: TensorPayload, -) -> DecodedSparsePayload: - """Flatten decoded locations so workers mmap only a few tensor storages.""" - packed_locations, value_groups, source_metadata = payload - location_parts: tuple[list[torch.Tensor], list[torch.Tensor]] = ([], []) - location_offsets = [0, 0] - metadata = [] - for item in source_metadata: - location_dtype = ( - torch.int32 - if prod(item["shape"]) <= torch.iinfo(torch.int32).max - else torch.int64 - ) - locations = sparse_locations_for_item( - item, packed_locations, device="cpu", dtype=location_dtype - ) - group = 0 if locations.dtype == torch.int32 else 1 - staged_item = dict(item) - staged_item["decoded_location_group"] = group - staged_item["decoded_location_start"] = location_offsets[group] - location_offsets[group] += locations.numel() - staged_item["decoded_location_end"] = location_offsets[group] - location_parts[group].append(locations) - metadata.append(staged_item) - locations = ( - _merge_tensor_parts(location_parts[0], torch.int32), - _merge_tensor_parts(location_parts[1], torch.int64), - ) - return locations, value_groups, metadata - - -def iter_decoded_sparse_payload( - payload: DecodedSparsePayload, -) -> Iterable[DecodedSparseItem]: - location_groups, value_groups, metadata = payload - for item in metadata: - location_start = int(item["decoded_location_start"]) - location_end = int(item["decoded_location_end"]) - value_start = int(item["value_start"]) - value_end = int(item["value_end"]) - yield ( - item, - location_groups[int(item["decoded_location_group"])][ - location_start:location_end - ], - value_groups[int(item["value_group"])][value_start:value_end], - ) - - def _encode_explicit_locations( locations: torch.Tensor, ) -> torch.Tensor: @@ -311,13 +217,12 @@ class DeltaCompressionTracker: def __init__( self, config: Mapping[str, Any], - *, - projections: Mapping[str, SparseShardProjection] | None = None, ) -> None: self.sparse_bucket_size_bytes = int(config["sparse_bucket_size_bytes"]) if self.sparse_bucket_size_bytes < 1: raise ValueError("delta_compression.sparse_bucket_size_bytes must be >= 1") self.encoding = sparse_operation(config["encoding"]) + self.overwrite_names: frozenset[str] = frozenset() self.verification_samples = int( os.getenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "0") ) @@ -325,7 +230,6 @@ def __init__( raise ValueError("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD must be >= 0") self.baseline_in_memory = os.getenv("NRL_REFIT_BASELINE_IN_MEMORY") == "1" self.baseline_mmap_dir = os.getenv("NRL_REFIT_BASELINE_MMAP_DIR") - self.projections = dict(projections or {}) self.baseline: dict[str, torch.Tensor] = {} self._pending_updates: dict[str, tuple[torch.Tensor, torch.Tensor]] = {} self._pending_updates_lock = threading.Lock() @@ -346,31 +250,21 @@ def prepare_sparse_delta_payload( total_elements += current.numel() changed_elements += locations.numel() if locations.numel(): + operation: SparseOperation = ( + "overwrite" if name in self.overwrite_names else self.encoding + ) values = ( current_values.bitwise_xor(baseline_bits[locations]) - if self.encoding == "xor" + if operation == "xor" else current_values ) - payload_name, payload_tensor, payload_locations = ( - name, - current, - locations, - ) - if projection := self.projections.get(name): - payload_name = projection.name - payload_locations = projection.map_locations( - locations, tuple(current.shape) - ) - payload_tensor = torch.empty( - projection.global_shape, dtype=current.dtype, device="meta" - ) sparse_infos.append( ( - payload_name, - payload_tensor, - payload_locations, + name, + current, + locations, values, - self.encoding, + operation, ) ) payload = encode_sparse_infos(sparse_infos) @@ -378,19 +272,6 @@ def prepare_sparse_delta_payload( self._add_verification_samples(payload[2]) return payload, changed_elements, total_elements - def prepare_change_summary( - self, tensors: Iterable[NamedTensor] - ) -> tuple[set[str], int, int]: - """Scan local tensors without constructing a wire payload.""" - changed_names = set() - changed_elements = total_elements = 0 - for name, _, current, locations, _ in self._changes(tensors): - total_elements += current.numel() - changed_elements += locations.numel() - if locations.numel(): - changed_names.add(name) - return changed_names, changed_elements, total_elements - def _changes( self, tensors: Iterable[NamedTensor] ) -> Iterable[tuple[str, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]]: @@ -483,7 +364,7 @@ def _baseline( if name in self.baseline: return self.baseline[name] numel = torch.Size(shape).numel() - nbytes = numel * torch.empty((), dtype=dtype).element_size() + nbytes = numel * dtype.itemsize if self.baseline_in_memory: storage = torch.empty(nbytes, dtype=torch.uint8) else: diff --git a/nemo_rl/utils/weight_transfer_remote_sparse.py b/nemo_rl/utils/weight_transfer_stream.py similarity index 71% rename from nemo_rl/utils/weight_transfer_remote_sparse.py rename to nemo_rl/utils/weight_transfer_stream.py index 8b6de5930fc..3bcd00f1839 100644 --- a/nemo_rl/utils/weight_transfer_remote_sparse.py +++ b/nemo_rl/utils/weight_transfer_stream.py @@ -12,27 +12,32 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Shared sparse payload pipeline and control plane for remote vLLM refit.""" +"""Shared sparse payload pipeline and S3 transport for remote vLLM refit.""" import hashlib import io import os import threading import time -from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence +from collections.abc import Iterable, Iterator, Mapping, Sequence from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait from contextlib import suppress from dataclasses import dataclass from functools import cache -from typing import Any, Literal +from typing import Any, Protocol from urllib.parse import quote -import requests import torch import zstandard -from urllib3.util.retry import Retry from nemo_rl.utils.packed_tensor import get_target_packed_tensor_size +from nemo_rl.utils.weight_transfer_http import ( + G_VLLM_REFIT_S3_MANIFEST_PATH, + merge_vllm_refit_metrics, + post_vllm_refit_endpoints, + vllm_refit_api_key, + vllm_refit_endpoints, +) from nemo_rl.utils.weight_transfer_sparse_codec import ( DeltaCompressionTracker, NamedTensor, @@ -41,16 +46,24 @@ merge_sparse_payloads, ) -G_VLLM_REFIT_S3_MANIFEST_PATH = "/nemo-rl/refit/s3-manifest" -G_VLLM_REFIT_PREPARE_PATH = "/nemo-rl/refit/prepare" -G_VLLM_REFIT_FLUSH_PATH = "/nemo-rl/refit/flush" -G_VLLM_REFIT_API_KEY_HEADER = "x-nemo-rl-refit-key" -_CONTROL_SESSION_LOCAL = threading.local() +_STREAM_LOCAL = threading.local() _S3_PART_SIZE = 64 * 1024**2 _S3_MEMORY_LIMIT = 2 * 1024**3 -SparsePartitionMode = Literal["chunks", "names", "none"] +class SparseRefitTransport(Protocol): + """Payload delivery owned by one generic stream invocation.""" + + name: str + transfer_workers: int + + def send( + self, body: bytes, payload_id: int, verification_candidates: int + ) -> dict[str, Any]: ... + + def cleanup(self) -> None: + """Release resources created on the current transfer worker.""" + ... @dataclass @@ -157,6 +170,75 @@ def _request(self, method: str, key: str, body: bytes | None = None) -> Any: ) +class _S3ManifestTransport: + name = "s3" + + def __init__( + self, + *, + store: _S3ObjectStore, + endpoint_urls: Sequence[str], + run_prefix: str, + api_key: str | None, + timeout_s: float, + transfer_workers: int, + ) -> None: + self._store = store + self._endpoint_urls = endpoint_urls + self._run_prefix = run_prefix + self._api_key = api_key + self._timeout_s = timeout_s + self.transfer_workers = transfer_workers + self._keys = threading.local() + + def send( + self, body: bytes, payload_id: int, verification_candidates: int + ) -> dict[str, Any]: + key = f"{self._run_prefix}/{payload_id:06d}.pt" + keys = getattr(self._keys, "values", None) + if keys is None: + keys = [] + self._keys.values = keys + keys.append(key) + + try: + started = time.perf_counter() + self._store.put(key, body) + s3_put_s = time.perf_counter() - started + started = time.perf_counter() + responses = post_vllm_refit_endpoints( + self._endpoint_urls, + { + "bucket": self._store.bucket, + "region": self._store.region, + "key": key, + "checksum": sparse_payload_checksum(body), + "verification_candidates": verification_candidates, + }, + api_key=self._api_key, + timeout_s=self._timeout_s, + ) + result = { + "s3_put_s": s3_put_s, + "manifest_post_s": time.perf_counter() - started, + "receiver": merge_vllm_refit_metrics({}, responses, maximum=True), + } + finally: + try: + self._store.delete(key) + except Exception: + pass + else: + keys.remove(key) + return result + + def cleanup(self) -> None: + for key in getattr(self._keys, "values", ()): + with suppress(Exception): + self._store.delete(key) + self._keys.values = [] + + def refit_env_int(name: str, *, default: int, min_value: int = 1) -> int: value = int(os.getenv(name) or default) if value < min_value: @@ -174,10 +256,10 @@ def decode_sparse_payload(body: bytes | bytearray, checksum: str) -> bytes: raise ValueError( f"Sparse refit payload checksum mismatch: expected={checksum}, actual={actual}." ) - decompressor = getattr(_CONTROL_SESSION_LOCAL, "zstd_decompressor", None) + decompressor = getattr(_STREAM_LOCAL, "zstd_decompressor", None) if decompressor is None: decompressor = zstandard.ZstdDecompressor() - _CONTROL_SESSION_LOCAL.zstd_decompressor = decompressor + _STREAM_LOCAL.zstd_decompressor = decompressor return decompressor.decompress(body) @@ -206,75 +288,19 @@ def iter_sparse_weight_chunks( yield chunk, export_pull_s -def sparse_name_shard(name: str, shard_count: int) -> int: - return ( - int.from_bytes(hashlib.blake2b(name.encode(), digest_size=8).digest()) - % shard_count - ) - - -def vllm_refit_endpoints(base_urls: Sequence[str], path: str) -> list[str]: - return list( - dict.fromkeys( - f"{url.strip().rstrip('/')}{path}" for url in base_urls if url.strip() - ) - ) - - -def _partition_sparse_weights_by_name( - tensors: Iterable[NamedTensor], shard_rank: int, shard_count: int -) -> Iterator[NamedTensor]: - return ( - (name, tensor) - for name, tensor in tensors - if sparse_name_shard(name, shard_count) == shard_rank - ) - - -def refit_http_session() -> requests.Session: - session = getattr(_CONTROL_SESSION_LOCAL, "session", None) - if session is None: - session = requests.Session() - adapter = requests.adapters.HTTPAdapter( - pool_connections=64, - pool_maxsize=64, - max_retries=Retry( - total=3, - backoff_factor=0.25, - status_forcelist=(502, 503, 504), - allowed_methods={"POST"}, - ), - ) - session.mount("http://", adapter) - session.mount("https://", adapter) - _CONTROL_SESSION_LOCAL.session = session - return session - - @cache def _get_manifest_s3_store(bucket: str, region: str) -> _S3ObjectStore: return _S3ObjectStore(bucket=bucket, region=region) -def vllm_refit_api_key(api_key_env_var: str | None) -> str | None: - if not api_key_env_var: - return None - token = os.environ.get(api_key_env_var) - if not token: - raise RuntimeError( - "vLLM sparse refit API key env var " - f"{api_key_env_var!r} is configured but unset or empty." - ) - return token - - def sparse_export_chunk_size( delta_tracker: DeltaCompressionTracker, transport: str, ) -> int: + default_mib = 64 if transport == "s3" else 256 requested = refit_env_int( f"NRL_REFIT_{transport.upper()}_EXPORT_CHUNK_BYTES", - default=256 * 1024**2, + default=default_mib * 1024**2, min_value=1, ) if torch.cuda.is_available(): @@ -294,10 +320,7 @@ def init_sparse_delta_baseline_from_iterator( shard_rank: int, shard_count: int, transport: str, - partition: SparsePartitionMode = "chunks", ) -> None: - if partition == "names": - iterator = _partition_sparse_weights_by_name(iterator, shard_rank, shard_count) start_s = time.perf_counter() export_chunk_size = sparse_export_chunk_size(delta_tracker, transport) @@ -308,7 +331,7 @@ def init_sparse_delta_baseline_from_iterator( ): chunk_count = chunk_index + 1 export_pull_s += pull_s - if partition == "chunks" and chunk_index % shard_count != shard_rank: + if chunk_index % shard_count != shard_rank: continue started = time.perf_counter() delta_tracker.snapshot_baseline(chunk) @@ -326,25 +349,25 @@ def stream_sparse_delta_payloads( iterator: Iterable[NamedTensor], *, delta_tracker: DeltaCompressionTracker, - transport: str, - send_payload: Callable[[bytes, int], dict[str, Any]], - transfer_workers: int, + transport: SparseRefitTransport, shard_rank: int, shard_count: int, - partition: SparsePartitionMode = "chunks", ) -> dict[str, int]: - if partition == "names": - iterator = _partition_sparse_weights_by_name(iterator, shard_rank, shard_count) - prefix = transport.upper() + prefix = transport.name.upper() encode_workers = refit_env_int( f"NRL_REFIT_{prefix}_ENCODE_WORKERS", default=max(2, min(8, os.cpu_count() or 8)), ) - encode_executor = _executor(f"refit-{transport}-encode", encode_workers) + encode_executor = _executor(f"refit-{transport.name}-encode", encode_workers) serialize_workers = min(4, encode_workers) - serialize_executor = _executor(f"refit-{transport}-serialize", serialize_workers) - transfer_executor = _executor(f"refit-{transport}-transfer", transfer_workers) - export_chunk_size = sparse_export_chunk_size(delta_tracker, transport) + serialize_executor = _executor( + f"refit-{transport.name}-serialize", serialize_workers + ) + transfer_executor = ThreadPoolExecutor( + max_workers=transport.transfer_workers, + thread_name_prefix=f"nrl-refit-{transport.name}-transfer", + ) + export_chunk_size = sparse_export_chunk_size(delta_tracker, transport.name) def encode_chunk( chunk: TensorBatch, @@ -365,10 +388,11 @@ def encode_chunk( def serialize_payloads( payloads: tuple[TensorPayload, ...], encode_s: float - ) -> tuple[bytes, dict[str, float]]: + ) -> tuple[bytes, int, dict[str, float]]: started = time.perf_counter() buffer = io.BytesIO() - torch.save(merge_sparse_payloads(payloads), buffer) + merged = merge_sparse_payloads(payloads) + torch.save(merged, buffer) raw_body = buffer.getvalue() serialize_s = time.perf_counter() - started started = time.perf_counter() @@ -376,6 +400,7 @@ def serialize_payloads( compress_s = time.perf_counter() - started return ( body, + sum(int(item.get("verification_samples", 0)) for item in merged[2]), { "encode_s": encode_s, "serialize_s": serialize_s, @@ -384,10 +409,10 @@ def serialize_payloads( ) def transfer_payload( - encoded: tuple[bytes, dict[str, float]], payload_index: int + encoded: tuple[bytes, int, dict[str, float]], payload_index: int ) -> dict[str, Any]: - body, encode_timing = encoded - result = send_payload(body, payload_index) + body, verification_candidates, encode_timing = encoded + result = transport.send(body, payload_index, verification_candidates) result.update( body_size=len(body), **encode_timing, @@ -407,9 +432,10 @@ def transfer_payload( encode_inflight: dict[Any, None] = {} serialize_inflight: dict[Any, int] = {} transfer_inflight: dict[Any, None] = {} + transfer_submitted = False max_encode_inflight = encode_workers * 2 max_serialize_inflight = serialize_workers * 2 - max_transfer_inflight = transfer_workers * 2 + max_transfer_inflight = transport.transfer_workers * 2 bucket = _SparsePayloadBucket([]) def resolved(inflight: dict[Any, Any], *, block: bool) -> Iterator[tuple[Any, Any]]: @@ -436,9 +462,11 @@ def collect_transfers(*, block: bool) -> None: ) def collect_serialized(*, block: bool) -> None: + nonlocal transfer_submitted for index, encoded in resolved(serialize_inflight, block=block): while len(transfer_inflight) >= max_transfer_inflight: collect_transfers(block=True) + transfer_submitted = True transfer_inflight[ transfer_executor.submit(transfer_payload, encoded, index) ] = None @@ -489,7 +517,7 @@ def drain_encodes() -> None: ): chunk_count = chunk_index + 1 export_pull_s += pull_s - if partition == "chunks" and chunk_index % shard_count != shard_rank: + if chunk_index % shard_count != shard_rank: continue while len(encode_inflight) >= max_encode_inflight: drain_encodes() @@ -509,6 +537,22 @@ def drain_encodes() -> None: if futures: wait(futures) raise + finally: + try: + if transfer_submitted: + barrier = threading.Barrier(transport.transfer_workers) + + def cleanup_transport(_index: int) -> None: + barrier.wait() + transport.cleanup() + + list( + transfer_executor.map( + cleanup_transport, range(transport.transfer_workers) + ) + ) + finally: + transfer_executor.shutdown(wait=True, cancel_futures=True) report = { "total_s": time.perf_counter() - stream_start, @@ -521,7 +565,6 @@ def drain_encodes() -> None: "export_chunk_mb": export_chunk_size / 1e6, "shard_rank": shard_rank, "shard_count": shard_count, - "partition": partition, "changed_elements": counts["changed_elements"], "total_elements": counts["total_elements"], "changed_pct": 100.0 @@ -551,10 +594,9 @@ def stream_sparse_delta_payloads_via_s3_manifest( timeout_s: float, shard_rank: int, shard_count: int, - partition: SparsePartitionMode = "chunks", ) -> dict[str, int]: - urls = [url.strip().rstrip("/") for url in refit_targets if url.strip()] - if not urls: + endpoint_urls = vllm_refit_endpoints(refit_targets, G_VLLM_REFIT_S3_MANIFEST_PATH) + if not endpoint_urls: raise ValueError("At least one vLLM S3 refit URL is required.") bucket = os.getenv("NRL_REFIT_S3_BUCKET", "").strip() if not bucket: @@ -563,89 +605,30 @@ def stream_sparse_delta_payloads_via_s3_manifest( bucket, os.getenv("NRL_REFIT_S3_REGION", "us-east-1").strip() or "us-east-1", ) - endpoint_urls = vllm_refit_endpoints(urls, G_VLLM_REFIT_S3_MANIFEST_PATH) object_prefix = os.getenv("NRL_REFIT_S3_PREFIX", "nemo-rl-refit").strip("/") run_prefix = ( f"{object_prefix}/{transfer_id}/{shard_rank:06d}" if object_prefix else f"{transfer_id}/{shard_rank:06d}" ) - api_key = vllm_refit_api_key(api_key_env_var) - - def send_payload(body: bytes, payload_index: int) -> dict[str, Any]: - key = f"{run_prefix}/{payload_index:06d}.pt" - started = time.perf_counter() - store.put(key, body) - s3_put_s = time.perf_counter() - started - try: - started = time.perf_counter() - responses = post_vllm_refit_endpoints( - endpoint_urls, - { - "bucket": store.bucket, - "region": store.region, - "key": key, - "checksum": sparse_payload_checksum(body), - }, - api_key=api_key, - timeout_s=timeout_s, - ) - manifest_post_s = time.perf_counter() - started - finally: - with suppress(Exception): - store.delete(key) - return { - "s3_put_s": s3_put_s, - "manifest_post_s": manifest_post_s, - "receiver": merge_vllm_refit_metrics({}, responses, maximum=True), - } - return stream_sparse_delta_payloads( iterator, delta_tracker=delta_tracker, - transport="s3", - send_payload=send_payload, - transfer_workers=refit_env_int( - "NRL_REFIT_S3_UPLOAD_WORKERS", - default=max(4, min(32, os.cpu_count() or 32)), + transport=_S3ManifestTransport( + store=store, + endpoint_urls=endpoint_urls, + run_prefix=run_prefix, + api_key=vllm_refit_api_key(api_key_env_var), + timeout_s=timeout_s, + transfer_workers=refit_env_int( + "NRL_REFIT_S3_UPLOAD_WORKERS", + default=max(4, min(32, os.cpu_count() or 32)), + ), ), shard_rank=shard_rank, shard_count=shard_count, - partition=partition, - ) - - -def post_vllm_refit_endpoints( - endpoint_urls: Sequence[str], - body: Mapping[str, Any] | bytes, - *, - api_key: str | None, - timeout_s: float, - headers: Mapping[str, str] | None = None, - executor: ThreadPoolExecutor | None = None, -) -> list[dict[str, Any]]: - request_headers = dict(headers or {}) - if api_key: - request_headers[G_VLLM_REFIT_API_KEY_HEADER] = api_key - request_kwargs: dict[str, Any] = ( - {"data": body} if isinstance(body, bytes) else {"json": body} ) - def post(url: str) -> dict[str, Any]: - response = refit_http_session().post( - url, - **request_kwargs, - headers=request_headers, - timeout=timeout_s, - ) - result: dict[str, Any] = response.json() if response.content else {} - if response.status_code >= 400 or result.get("ok") is not True: - raise RuntimeError(f"vLLM refit failed for {url}: {result}") - return result - - pool = executor or _executor("refit-fanout", len(endpoint_urls)) - return list(pool.map(post, endpoint_urls)) - def download_s3_refit_payload( manifest: Mapping[str, Any], @@ -656,38 +639,12 @@ def download_s3_refit_payload( return decode_sparse_payload(body, str(manifest["checksum"])) -def merge_vllm_refit_metrics( - result: dict[str, Any], - metrics: Iterable[Mapping[str, Any]], - *, - maximum: bool, - candidate_maximum: bool | None = None, -) -> dict[str, Any]: - for metric in metrics: - for key, value in metric.items(): - if key.startswith("receiver_") and key.endswith("_s"): - number, use_maximum = float(value), maximum - elif candidate_maximum is not None and key.startswith("verification_"): - number = value - use_maximum = key == "verification_max_abs" or ( - key == "verification_candidates" and candidate_maximum - ) - else: - continue - if key in result: - number = ( - max(result[key], number) if use_maximum else result[key] + number - ) - result[key] = number - return result - - def zstd_compress(raw: bytes, threads_env: str) -> bytes: - compressor = getattr(_CONTROL_SESSION_LOCAL, "zstd_compressor", None) + compressor = getattr(_STREAM_LOCAL, "zstd_compressor", None) if compressor is None: compressor = zstandard.ZstdCompressor( level=1, threads=refit_env_int(threads_env, default=0, min_value=0), ) - _CONTROL_SESSION_LOCAL.zstd_compressor = compressor + _STREAM_LOCAL.zstd_compressor = compressor return compressor.compress(raw) diff --git a/nemo_rl/utils/weight_transfer_zmq.py b/nemo_rl/utils/weight_transfer_zmq.py index fd3020dde1d..fb00b998955 100644 --- a/nemo_rl/utils/weight_transfer_zmq.py +++ b/nemo_rl/utils/weight_transfer_zmq.py @@ -19,19 +19,16 @@ import time import uuid from collections.abc import Iterable, Mapping, Sequence -from concurrent.futures import ThreadPoolExecutor +from concurrent.futures import Future, ThreadPoolExecutor from contextlib import suppress +from dataclasses import dataclass, field from typing import Any import zmq -from nemo_rl.utils.weight_transfer_remote_sparse import ( - SparsePartitionMode, +from nemo_rl.utils.weight_transfer_http import ( merge_vllm_refit_metrics, post_vllm_refit_endpoints, - refit_env_int, - sparse_payload_checksum, - stream_sparse_delta_payloads, vllm_refit_api_key, vllm_refit_endpoints, ) @@ -39,18 +36,29 @@ DeltaCompressionTracker, NamedTensor, ) +from nemo_rl.utils.weight_transfer_stream import ( + refit_env_int, + sparse_payload_checksum, + stream_sparse_delta_payloads, +) G_VLLM_REFIT_ZMQ_PAYLOAD_PATH = "/nemo-rl/refit/zmq-payload" G_VLLM_REFIT_TRANSFER_HEADER = "x-nemo-rl-refit-transfer" G_VLLM_REFIT_PRODUCER_HEADER = "x-nemo-rl-refit-producer" G_VLLM_REFIT_PAYLOAD_HEADER = "x-nemo-rl-refit-payload" G_VLLM_REFIT_CHECKSUM_HEADER = "x-nemo-rl-refit-checksum" +G_VLLM_REFIT_VERIFICATION_HEADER = "x-nemo-rl-refit-verification-candidates" _PROTOCOL = "nemo-rl-sparse-zmq-v1" _DATA = b"DATA" _ACK = b"ACK" _NACK = b"NACK" -_ZMQ_LOCAL = threading.local() + + +@dataclass +class _RelayTransfer: + checksums: dict[tuple[int, int], str] = field(default_factory=dict) + futures: list[Future[dict[str, Any]]] = field(default_factory=list) def _json_bytes(value: Mapping[str, Any]) -> bytes: @@ -94,15 +102,22 @@ def send_payload( transfer_id: str, payload_id: int, checksum: str, + verification_candidates: int, body: bytes, + relay_root: str | None = None, + producer_id: int | None = None, ) -> dict[str, Any]: + producer_id = self._producer_id if producer_id is None else producer_id metadata = { "protocol": _PROTOCOL, "transfer_id": transfer_id, - "producer_id": self._producer_id, + "producer_id": producer_id, "payload_id": payload_id, "checksum": checksum, + "verification_candidates": verification_candidates, } + if relay_root is not None: + metadata["relay_root"] = relay_root if self._api_key is not None: metadata["api_key"] = self._api_key metadata_frame = _json_bytes(metadata) @@ -138,7 +153,7 @@ def send_payload( reply.get("producer_id"), reply.get("payload_id"), ) - if reply_key != (transfer_id, self._producer_id, payload_id): + if reply_key != (transfer_id, producer_id, payload_id): continue if kind == _ACK and reply.get("ok") is True: return reply @@ -155,6 +170,59 @@ def close(self) -> None: self._socket.close() +class _ZmqTransport: + name = "zmq" + + def __init__( + self, + address: str, + *, + transfer_id: str, + timeout_s: float, + producer_id: int, + api_key: str | None, + transfer_workers: int, + ) -> None: + self._address = address + self._transfer_id = transfer_id + self._timeout_s = timeout_s + self._producer_id = producer_id + self._api_key = api_key + self.transfer_workers = transfer_workers + self._local = threading.local() + + def send( + self, body: bytes, payload_id: int, verification_candidates: int + ) -> dict[str, Any]: + client = getattr(self._local, "client", None) + if client is None: + client = ZmqSparseRefitClient( + self._address, + timeout_s=self._timeout_s, + producer_id=self._producer_id, + api_key=self._api_key, + ) + self._local.client = client + started = time.perf_counter() + reply = client.send_payload( + transfer_id=self._transfer_id, + payload_id=payload_id, + checksum=sparse_payload_checksum(body), + verification_candidates=verification_candidates, + body=body, + ) + return { + "zmq_send_s": time.perf_counter() - started, + "receiver": reply, + } + + def cleanup(self) -> None: + client = getattr(self._local, "client", None) + if client is not None: + client.close() + del self._local.client + + class ZmqSparseRefitServer: """Bounded ROUTER relay that fans each compressed payload to all replicas.""" @@ -184,9 +252,41 @@ def __init__( ) self._fanout_workers = refit_env_int( "NRL_REFIT_ZMQ_RELAY_FANOUT_WORKERS", - default=max(8, min(32, len(self._refit_endpoints) * self._payload_workers)), + default=len(self._refit_endpoints) * self._payload_workers, + ) + self._transfer_lock = threading.Lock() + self._transfer_condition = threading.Condition(self._transfer_lock) + self._transfers: dict[str, _RelayTransfer] = {} + self._flush_results: dict[str, Future[dict[str, Any]]] = {} + self._tree: tuple[tuple[str, ...], int, str] | None = None + self._forward_local = threading.local() + self._forward_clients: list[ZmqSparseRefitClient] = [] + self._payload_executor = ThreadPoolExecutor( + max_workers=self._payload_workers, + thread_name_prefix="nrl-zmq-payload", + ) + self._forward_executor = ThreadPoolExecutor( + max_workers=refit_env_int("NRL_REFIT_ZMQ_RELAY_FORWARD_WORKERS", default=8), + thread_name_prefix="nrl-zmq-forward", + ) + self._http_executor = ThreadPoolExecutor( + max_workers=self._fanout_workers, + thread_name_prefix="nrl-zmq-fanout", ) + def configure_tree( + self, + relay_addresses: Sequence[str], + *, + own_address: str, + local_refit_url: str, + ) -> None: + addresses = tuple(dict.fromkeys(relay_addresses)) + (local_endpoint,) = vllm_refit_endpoints( + [local_refit_url], G_VLLM_REFIT_ZMQ_PAYLOAD_PATH + ) + self._tree = (addresses, addresses.index(own_address), local_endpoint) + def start(self) -> str: self._thread = threading.Thread( target=self._run, @@ -211,11 +311,65 @@ def close(self) -> None: raise RuntimeError("Timed out stopping the ZeroMQ sparse refit relay.") self._thread = None + def flush(self, transfer_id: str, expected_payloads: int = 0) -> dict[str, Any]: + """Wait for every staged fanout belonging to one transfer.""" + with self._transfer_condition: + completion = self._flush_results.get(transfer_id) + if completion is None: + ready = self._transfer_condition.wait_for( + lambda: len( + self._transfers.get(transfer_id, _RelayTransfer()).checksums + ) + >= expected_payloads, + timeout=self._timeout_s, + ) + if not ready: + raise TimeoutError( + f"Timed out waiting for {expected_payloads} ZeroMQ payloads " + f"for transfer {transfer_id}." + ) + completion = self._flush_results.get(transfer_id) + if completion is None: + completion = Future() + self._flush_results[transfer_id] = completion + staged = self._transfers.pop(transfer_id, _RelayTransfer()) + else: + staged = None + if staged is None: + return completion.result() + + try: + started = time.perf_counter() + results: list[dict[str, Any]] = [] + first_error: Exception | None = None + for future in staged.futures: + try: + results.append(future.result()) + except Exception as exc: + if first_error is None: + first_error = exc + if first_error is not None: + raise RuntimeError( + f"ZeroMQ relay fanout failed for transfer {transfer_id}: " + f"{first_error}" + ) from first_error + + merged = merge_vllm_refit_metrics({}, results, maximum=False) + merged.update( + ok=True, + payloads=len(staged.checksums), + receiver_relay_flush_s=time.perf_counter() - started, + ) + except Exception as exc: + completion.set_exception(exc) + raise + completion.set_result(merged) + return merged + def _fanout( self, body: bytes, metadata: Mapping[str, Any], - http_executor: ThreadPoolExecutor, ) -> dict[str, Any]: headers = { "content-type": "application/octet-stream", @@ -223,20 +377,58 @@ def _fanout( G_VLLM_REFIT_PRODUCER_HEADER: str(metadata["producer_id"]), G_VLLM_REFIT_PAYLOAD_HEADER: str(metadata["payload_id"]), G_VLLM_REFIT_CHECKSUM_HEADER: str(metadata["checksum"]), + G_VLLM_REFIT_VERIFICATION_HEADER: str(metadata["verification_candidates"]), } started = time.perf_counter() + endpoints = self._refit_endpoints + if self._tree is not None: + endpoints = [self._tree[2]] results = post_vllm_refit_endpoints( - self._refit_endpoints, + endpoints, body, api_key=self._token, timeout_s=self._timeout_s, headers=headers, - executor=http_executor, + executor=self._http_executor, ) merged = merge_vllm_refit_metrics({}, results, maximum=True) merged["receiver_relay_fanout_s"] = time.perf_counter() - started return merged + def _forward( + self, + body: bytes, + metadata: Mapping[str, Any], + address: str, + relay_root: str, + ) -> dict[str, Any]: + clients = getattr(self._forward_local, "clients", None) + if clients is None: + clients = {} + self._forward_local.clients = clients + client = clients.get(address) + if client is None: + client = ZmqSparseRefitClient( + address, + timeout_s=self._timeout_s, + producer_id=0, + api_key=self._token, + ) + clients[address] = client + with self._transfer_lock: + self._forward_clients.append(client) + started = time.perf_counter() + client.send_payload( + transfer_id=str(metadata["transfer_id"]), + payload_id=int(metadata["payload_id"]), + checksum=str(metadata["checksum"]), + verification_candidates=int(metadata["verification_candidates"]), + body=body, + relay_root=relay_root, + producer_id=int(metadata["producer_id"]), + ) + return {"receiver_relay_forward_s": time.perf_counter() - started} + @staticmethod def _send_reply( socket: zmq.Socket, @@ -269,20 +461,16 @@ def _parse_data_message( checksum = str(metadata["checksum"]) if not transfer_id or producer_id < 0 or payload_id < 0: raise ValueError("Invalid ZeroMQ sparse refit payload identity.") + relay_root = metadata.get("relay_root") + if relay_root is not None and ( + self._tree is None or relay_root not in self._tree[0] + ): + raise ValueError("Invalid ZeroMQ relay root.") return identity, (transfer_id, producer_id, payload_id), body, metadata def _run(self) -> None: context = zmq.Context() socket = context.socket(zmq.ROUTER) - payload_executor = ThreadPoolExecutor( - max_workers=self._payload_workers, - thread_name_prefix="nrl-zmq-payload", - ) - http_executor = ThreadPoolExecutor( - max_workers=self._fanout_workers, - thread_name_prefix="nrl-zmq-fanout", - ) - pending: dict[Any, tuple[bytes, tuple[str, int, int]]] = {} try: _configure_socket(socket, 16) socket.setsockopt(zmq.ROUTER_MANDATORY, 1) @@ -290,23 +478,75 @@ def _run(self) -> None: self._endpoint = socket.getsockopt_string(zmq.LAST_ENDPOINT) self._ready.set() - while not self._stop.is_set() or pending: - if ( - not self._stop.is_set() - and len(pending) < self._payload_workers - and socket.poll(10, zmq.POLLIN) - ): + while not self._stop.is_set(): + if socket.poll(10, zmq.POLLIN): frames = socket.recv_multipart() identity = frames[0] if frames else b"" try: identity, key, body, metadata = self._parse_data_message(frames) - future = payload_executor.submit( - self._fanout, - body, - metadata, - http_executor, + transfer_id, producer_id, payload_id = key + payload_key = (producer_id, payload_id) + checksum = str(metadata["checksum"]) + with self._transfer_lock: + if transfer_id in self._flush_results: + raise RuntimeError( + "ZeroMQ sparse refit transfer is already flushed." + ) + staged = self._transfers.setdefault( + transfer_id, _RelayTransfer() + ) + previous = staged.checksums.get(payload_key) + if previous is not None and previous != checksum: + raise ValueError( + "Conflicting ZeroMQ sparse refit payload checksum." + ) + if previous is None: + staged.checksums[payload_key] = checksum + staged.futures.append( + self._payload_executor.submit( + self._fanout, body, metadata + ) + ) + if self._tree is not None: + addresses, own_index, _ = self._tree + relay_root = str( + metadata.get("relay_root", addresses[own_index]) + ) + root_index = addresses.index(relay_root) + node_index = (own_index - root_index) % len( + addresses + ) + for child_index in ( + 2 * node_index + 1, + 2 * node_index + 2, + ): + if child_index < len(addresses): + child = addresses[ + (root_index + child_index) + % len(addresses) + ] + staged.futures.append( + self._forward_executor.submit( + self._forward, + body, + metadata, + child, + relay_root, + ) + ) + self._transfer_condition.notify_all() + self._send_reply( + socket, + identity, + _ACK, + { + "ok": True, + "staged": True, + "transfer_id": transfer_id, + "producer_id": producer_id, + "payload_id": payload_id, + }, ) - pending[future] = (identity, key) except Exception as exc: self._send_reply( socket, @@ -314,32 +554,15 @@ def _run(self) -> None: _NACK, {"ok": False, "error": str(exc)}, ) - - for future, (identity, key) in list(pending.items()): - if not future.done(): - continue - transfer_id, producer_id, payload_id = key - reply: dict[str, Any] = { - "ok": True, - "transfer_id": transfer_id, - "producer_id": producer_id, - "payload_id": payload_id, - } - try: - reply.update(future.result()) - except Exception as exc: - reply.update(ok=False, error=str(exc)) - kind = _NACK - else: - kind = _ACK - self._send_reply(socket, identity, kind, reply) - del pending[future] except Exception as exc: self._error = exc self._ready.set() finally: - payload_executor.shutdown(wait=True, cancel_futures=True) - http_executor.shutdown(wait=True, cancel_futures=True) + self._payload_executor.shutdown(wait=True, cancel_futures=True) + self._forward_executor.shutdown(wait=True, cancel_futures=True) + self._http_executor.shutdown(wait=True, cancel_futures=True) + for client in self._forward_clients: + client.close() socket.close() context.term() @@ -354,48 +577,22 @@ def stream_sparse_delta_payloads_via_zmq( timeout_s: float, shard_rank: int, shard_count: int, - partition: SparsePartitionMode = "chunks", ) -> dict[str, int]: addresses = [address.strip() for address in refit_targets if address.strip()] if not addresses: raise ValueError("At least one ZeroMQ sparse refit address is required.") address = addresses[shard_rank % len(addresses)] - api_key = vllm_refit_api_key(api_key_env_var) - - def send_payload(body: bytes, payload_id: int) -> dict[str, Any]: - clients = getattr(_ZMQ_LOCAL, "clients", None) - if clients is None: - clients = {} - _ZMQ_LOCAL.clients = clients - client_key = (address, shard_rank, api_key) - client = clients.get(client_key) - if client is None: - client = ZmqSparseRefitClient( - address, - timeout_s=timeout_s, - producer_id=shard_rank, - api_key=api_key, - ) - clients[client_key] = client - started = time.perf_counter() - reply = client.send_payload( - transfer_id=transfer_id, - payload_id=payload_id, - checksum=sparse_payload_checksum(body), - body=body, - ) - return { - "zmq_send_s": time.perf_counter() - started, - "receiver": reply, - } - return stream_sparse_delta_payloads( iterator, delta_tracker=delta_tracker, - transport="zmq", - send_payload=send_payload, - transfer_workers=refit_env_int("NRL_REFIT_ZMQ_SEND_WORKERS", default=4), + transport=_ZmqTransport( + address, + transfer_id=transfer_id, + timeout_s=timeout_s, + producer_id=shard_rank, + api_key=vllm_refit_api_key(api_key_env_var), + transfer_workers=refit_env_int("NRL_REFIT_ZMQ_SEND_WORKERS", default=4), + ), shard_rank=shard_rank, shard_count=shard_count, - partition=partition, ) diff --git a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py index fa88fb94870..b55222bb600 100644 --- a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py +++ b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py @@ -22,10 +22,15 @@ import ray +from nemo_rl.models.generation.vllm.config import ( + VllmConfig, + VllmDeltaCompressionConfig, +) from nemo_rl.utils.timer import Timer -from nemo_rl.utils.weight_transfer_remote_sparse import ( +from nemo_rl.utils.weight_transfer_http import ( G_VLLM_REFIT_FLUSH_PATH, G_VLLM_REFIT_PREPARE_PATH, + G_VLLM_REFIT_ZMQ_FLUSH_PATH, merge_vllm_refit_metrics, post_vllm_refit_endpoints, vllm_refit_api_key, @@ -40,7 +45,7 @@ def validate_vllm_remote_sparse_refit( - config: Any, + config: VllmConfig, *, colocated: bool, megatron_enabled: bool, @@ -52,12 +57,13 @@ def validate_vllm_remote_sparse_refit( if transport not in _REMOTE_SPARSE_TRANSPORTS: raise ValueError(f"Unsupported vLLM refit transport {transport!r}.") vllm_cfg = config["vllm_cfg"] + delta_config = config.get("delta_compression") if ( colocated or not megatron_enabled or vllm_cfg["precision"] == "fp8" or vllm_cfg["kv_cache_dtype"].startswith("fp8") - or not config.get("delta_compression") + or delta_config is None or config.get("quant_cfg") or config.get("real_quant") ): @@ -65,6 +71,9 @@ def validate_vllm_remote_sparse_refit( f"{transport} requires a non-colocated Megatron policy, BF16/FP16 " "vLLM, delta compression, and an unquantized rollout." ) + config["delta_compression"] = VllmDeltaCompressionConfig.model_validate( + delta_config + ) return _REMOTE_SPARSE_TRANSPORTS[transport] @@ -86,6 +95,7 @@ def __init__( self._request_timeout_s = request_timeout_s self._refit_urls: list[str] = [] self._targets: list[str] = [] + self._overwrite_names: list[str] = [] self._baseline_init_refs = list(baseline_init_refs or ()) self._baseline_commit_refs: list[Any] = [] self._stale = True @@ -115,8 +125,10 @@ def sync_weights( self._baseline_init_refs.clear() succeeded = False + relay_flushed = False + relay_flush_s = 0.0 + transfer_id = uuid.uuid4().hex try: - transfer_id = uuid.uuid4().hex results = ray.get( self._run_policy_workers( "stream_remote_sparse_weights", @@ -125,6 +137,7 @@ def sync_weights( transfer_id=transfer_id, api_key_env_var=self._api_key_env_var, timeout_s=self._request_timeout_s, + overwrite_names=self._overwrite_names, ) ) payloads = sum(result["payloads"] for result in results) @@ -141,6 +154,32 @@ def sync_weights( verification: defaultdict[str, float] = defaultdict(float) commit_s = 0.0 if payloads: + if self._transport == "zmq": + started = time.perf_counter() + relay_results = self._request_receivers( + G_VLLM_REFIT_ZMQ_FLUSH_PATH, + { + "transfer_id": transfer_id, + "expected_payloads": payloads, + }, + ) + relay_flush_s = time.perf_counter() - started + relay_flushed = True + if any( + int(result.get("payloads", 0)) != payloads + for result in relay_results + ): + raise RuntimeError( + f"ZeroMQ relays did not all stage {payloads} payloads." + ) + print( + "REFIT_ZMQ_RELAY_FLUSH " + f"transfer_id={transfer_id} payloads={payloads} " + f"seconds={relay_flush_s:.3f} " + "fanout_service_s=" + f"{sum(float(result.get('receiver_relay_fanout_s', 0.0)) for result in relay_results):.3f}", + flush=True, + ) started = time.perf_counter() verification.update( merge_vllm_refit_metrics( @@ -180,6 +219,13 @@ def sync_weights( succeeded = True finally: if not succeeded: + if self._transport == "zmq" and not relay_flushed: + with suppress(Exception): + self._request_receivers( + G_VLLM_REFIT_ZMQ_FLUSH_PATH, + {"transfer_id": transfer_id}, + timeout_s=min(self._request_timeout_s, 60.0), + ) with suppress(Exception): self._request_receivers( G_VLLM_REFIT_FLUSH_PATH, @@ -203,6 +249,7 @@ def sync_weights( "delta_verify/mean_abs": float(verification["verification_abs_sum"]) / max(samples, 1), "transfer/payloads": float(payloads), + "transfer/relay_flush_s": relay_flush_s, "transfer/global_commit_s": commit_s, } metrics.update( @@ -302,11 +349,15 @@ def init_communicator(self) -> None: raise ValueError( f"vLLM {self._transport} sparse refit endpoints are missing." ) + if self._transport == "zmq": + self._run_generation_workers( + "configure_zmq_sparse_refit_relay", relay_addresses=self._targets + ) state_dict_info = self._merge_refit_info( ray.get(list(self._baseline_init_refs)) ) self._baseline_init_refs.clear() - self._request_receivers( + responses = self._request_receivers( G_VLLM_REFIT_PREPARE_PATH, { "tensors": { @@ -315,6 +366,13 @@ def init_communicator(self) -> None: } }, ) + self._overwrite_names = sorted( + { + name + for response in responses + for name in response.get("overwrite_names", ()) + } + ) self._stale = False def shutdown(self) -> None: @@ -326,4 +384,5 @@ def shutdown(self) -> None: self._baseline_commit_refs.clear() self._refit_urls.clear() self._targets.clear() + self._overwrite_names.clear() self._stale = True diff --git a/pyrefly.toml b/pyrefly.toml index 6458fc92148..75315f5d862 100644 --- a/pyrefly.toml +++ b/pyrefly.toml @@ -197,8 +197,9 @@ project-includes = [ "nemo_rl/utils/r3_trace.py", "nemo_rl/utils/timer.py", "nemo_rl/utils/venvs.py", - "nemo_rl/utils/weight_transfer_remote_sparse.py", + "nemo_rl/utils/weight_transfer_http.py", "nemo_rl/utils/weight_transfer_sparse_codec.py", + "nemo_rl/utils/weight_transfer_stream.py", "nemo_rl/utils/weight_transfer_zmq.py", "nemo_rl/weight_sync/__init__.py", "nemo_rl/weight_sync/collective_weight_synchronizer.py", diff --git a/tests/unit/models/generation/test_vllm_backend.py b/tests/unit/models/generation/test_vllm_backend.py index ea667f6c86c..2fd83278116 100644 --- a/tests/unit/models/generation/test_vllm_backend.py +++ b/tests/unit/models/generation/test_vllm_backend.py @@ -33,7 +33,7 @@ def _make_collective_update_extension(backend): state_info = object() ext.state_dict_info = {"model.weight": state_info} ext.model_update_group = object() - ext.model_runner = SimpleNamespace(model=object()) + ext.model_runner = SimpleNamespace(model=object(), vllm_config=object()) ext.model_config = object() ext.device = object() return ext, state_info diff --git a/tests/unit/models/generation/test_vllm_sparse_delta.py b/tests/unit/models/generation/test_vllm_sparse_delta.py index 2ea6c851cb9..76f0c47814b 100644 --- a/tests/unit/models/generation/test_vllm_sparse_delta.py +++ b/tests/unit/models/generation/test_vllm_sparse_delta.py @@ -26,9 +26,7 @@ ) from nemo_rl.utils.weight_transfer_sparse_codec import ( SparseOperation, - decode_sparse_tensor_payload_for_staging, encode_sparse_infos, - iter_decoded_sparse_payload, ) from nemo_rl.utils.weight_transfer_sparse_codec import ( integer_view as _bits, @@ -88,12 +86,6 @@ def _applier(model: Any) -> VllmSparseDeltaApplier: ) -def _decode_staged(payload: Any) -> list[Any]: - return list( - iter_decoded_sparse_payload(decode_sparse_tensor_payload_for_staging(payload)) - ) - - def _payload( name: str, tensor: torch.Tensor, @@ -107,21 +99,22 @@ def _payload( def _apply_payload(applier: VllmSparseDeltaApplier, payload: Any) -> None: - applier._apply_decoded_items(_decode_staged(payload)) + buffer = io.BytesIO() + torch.save(payload, buffer) + applier.update_weights_from_decoded_sparse_payload(buffer.getvalue()) -def test_sparse_prewarm_reserves_largest_source_without_loading_weights() -> None: +def test_sparse_discovery_reserves_largest_source_without_loading_weights() -> None: target = torch.zeros(8) applier = _applier(_NativeLoaderModel(identity=target)) - applier.prewarm({"weight": ((8,), torch.float32)}) + applier.discover_native_skips({"weight": ((8,), torch.float32)}) assert applier._scratch.numel() == target.numel() * target.element_size() assert torch.equal(target, torch.zeros_like(target)) -@pytest.mark.vllm -def test_sparse_prewarm_caches_rank_local_native_loader_skips() -> None: +def test_sparse_discovery_caches_rank_local_native_loader_skips() -> None: target = torch.zeros(2) model = _NativeLoaderModel(identity=target) applier = _applier(model) @@ -130,15 +123,31 @@ def test_sparse_prewarm_caches_rank_local_native_loader_skips() -> None: "skipped": ((2,), torch.float32), } - applier.prewarm(info) - applier.discover_native_skips(info) + overwrite_names = applier.discover_native_skips(info) payload = _payload("skipped", target, [1], _bits(torch.tensor([3.0]))) _apply_payload(applier, payload) assert model.loaded_names == ["weight", "skipped"] + assert overwrite_names == set() assert torch.equal(target, torch.zeros_like(target)) +def test_sparse_discovery_classifies_xor_unsafe_native_loads() -> None: + model = _NativeLoaderModel( + identity=torch.zeros(2), + exp_transform=torch.zeros(2), + ) + applier = _applier(model) + info = { + "weight": ((2,), torch.bfloat16), + "exp_transform": ((2,), torch.float32), + } + + assert applier.discover_native_skips(info) == {"weight", "exp_transform"} + assert torch.equal(model.targets["identity"], torch.zeros(2)) + assert torch.equal(model.targets["exp_transform"], torch.zeros(2)) + + @pytest.mark.vllm def test_backend_applies_decoded_sparse_payload_sources() -> None: from nemo_rl.models.generation.vllm.vllm_backend import ( @@ -147,10 +156,13 @@ def test_backend_applies_decoded_sparse_payload_sources() -> None: ext = VllmInternalWorkerExtension.__new__(VllmInternalWorkerExtension) applier = MagicMock() + applier.discover_native_skips.return_value = {"weight"} applier.update_weights_from_decoded_sparse_payload.return_value = {"ok": True} ext._get_sparse_delta_applier = MagicMock(return_value=applier) - ext.prepare_sparse_delta_refit_info({"weight": ((8,), torch.float32)}) + assert ext.prepare_sparse_delta_refit_info({"weight": ((8,), torch.float32)}) == [ + "weight" + ] assert ext.update_weights_from_decoded_sparse_payload(b"payload") == {"ok": True} assert ext.update_weights_from_decoded_sparse_payload("first", "second") == { @@ -160,38 +172,34 @@ def test_backend_applies_decoded_sparse_payload_sources() -> None: item.args for item in applier.update_weights_from_decoded_sparse_payload.call_args_list ] == [(b"payload",), ("first", "second")] - applier.prewarm.assert_called_once_with({"weight": ((8,), torch.float32)}) applier.discover_native_skips.assert_called_once_with( {"weight": ((8,), torch.float32)} ) -@pytest.mark.vllm def test_sparse_payload_batches_preserve_order(tmp_path) -> None: applier = _applier(_NativeLoaderModel(identity=torch.zeros(1))) - decoded_paths = [tmp_path / f"decoded-{index}.pt" for index in range(3)] - decoded_payloads = [ - decode_sparse_tensor_payload_for_staging( - _payload( - f"weight-{index}", - torch.empty(1), - [0], - _bits(torch.tensor([float(index)])), - ) + payload_paths = [tmp_path / f"payload-{index}.pt" for index in range(3)] + payloads = [ + _payload( + f"weight-{index}", + torch.empty(1), + [0], + _bits(torch.tensor([float(index)])), ) for index in range(3) ] - for path, payload in zip(decoded_paths, decoded_payloads, strict=True): + for path, payload in zip(payload_paths, payloads, strict=True): torch.save(payload, path) decoded_applied: list[Any] = [] applier._apply_decoded_items = lambda items: decoded_applied.extend( item for item, _, _ in items ) result = applier.update_weights_from_decoded_sparse_payload( - *(path.read_bytes() for path in decoded_paths) + *(path.read_bytes() for path in payload_paths) ) decoded_result = applier.update_weights_from_decoded_sparse_payload( - *(str(path) for path in reversed(decoded_paths)) + *(str(path) for path in reversed(payload_paths)) ) assert [item["name"] for item in decoded_applied] == [ @@ -207,7 +215,6 @@ def test_sparse_payload_batches_preserve_order(tmp_path) -> None: assert decoded_result["receiver_deserialize_s"] >= 0.0 -@pytest.mark.vllm def test_sparse_payload_batch_uses_one_streaming_native_loader_call() -> None: identity = torch.zeros(2) scale = torch.zeros(2) @@ -218,7 +225,7 @@ def test_sparse_payload_batch_uses_one_streaming_native_loader_call() -> None: _payload("weight_scale_inv", scale, [1], _bits(torch.tensor([3.0]))), ): buffer = io.BytesIO() - torch.save(decode_sparse_tensor_payload_for_staging(payload), buffer) + torch.save(payload, buffer) serialized.append(buffer.getvalue()) _applier(model).update_weights_from_decoded_sparse_payload(*serialized) @@ -228,14 +235,11 @@ def test_sparse_payload_batch_uses_one_streaming_native_loader_call() -> None: assert torch.equal(scale, torch.tensor([0.0, 3.0])) -@pytest.mark.vllm -def test_decoded_sparse_payload_converts_compact_locations_for_apply(tmp_path) -> None: +def test_compact_sparse_payload_decodes_locations_for_apply(tmp_path) -> None: target = torch.zeros(8) payload = _payload("weight", target, [1, 5], _bits(torch.tensor([2.0, 6.0]))) - decoded = decode_sparse_tensor_payload_for_staging(payload) - assert decoded[0][0].dtype == torch.int32 - path = tmp_path / "decoded.pt" - torch.save(decoded, path) + path = tmp_path / "payload.pt" + torch.save(payload, path) result = _applier( _NativeLoaderModel(identity=target) @@ -245,7 +249,6 @@ def test_decoded_sparse_payload_converts_compact_locations_for_apply(tmp_path) - assert result["receiver_sparse_apply_s"] >= 0.0 -@pytest.mark.vllm def test_native_loaders_apply_sparse_views_and_transforms() -> None: targets = { "identity": torch.zeros(4), @@ -301,7 +304,6 @@ def test_native_loaders_apply_sparse_views_and_transforms() -> None: assert verification["verification_exact_mismatches"] == 0 -@pytest.mark.vllm def test_xor_applies_through_packed_native_loaders() -> None: targets = { "row_slice": torch.zeros(2, 2), @@ -339,18 +341,15 @@ def test_xor_applies_through_packed_native_loaders() -> None: ] -@pytest.mark.vllm def test_native_loader_explicit_skip_is_accepted() -> None: model = _NativeLoaderModel(identity=torch.zeros(1)) payload = _payload("skipped", torch.empty(1), [0], _bits(torch.tensor([1.0]))) - item, locations, values = _decode_staged(payload)[0] - _applier(model)._apply_decoded_items(((item, locations, values),)) + _apply_payload(_applier(model), payload) assert torch.equal(model.targets["identity"], torch.zeros(1)) -@pytest.mark.vllm def test_native_loader_claim_without_copy_fails_closed() -> None: model = _NativeLoaderModel(identity=torch.zeros(1)) model.load_weights = lambda weights: {name for name, _ in weights} @@ -360,7 +359,6 @@ def test_native_loader_claim_without_copy_fails_closed() -> None: _apply_payload(_applier(model), payload) -@pytest.mark.vllm def test_sparse_overwrite_preserves_unselected_transform_inputs() -> None: target = torch.tensor([-2.0, -4.0, -6.0, -8.0]) payload = _payload( @@ -375,7 +373,6 @@ def test_sparse_overwrite_preserves_unselected_transform_inputs() -> None: assert torch.allclose(target, torch.tensor([-2.0, -3.0, -6.0, -8.0])) -@pytest.mark.vllm def test_unknown_sparse_operation_fails_closed() -> None: target = torch.zeros(1) payload = _payload("weight", target, [0], _bits(torch.tensor([1.0]))) @@ -385,7 +382,6 @@ def test_unknown_sparse_operation_fails_closed() -> None: _apply_payload(_applier(_NativeLoaderModel(identity=target)), payload) -@pytest.mark.vllm @pytest.mark.parametrize( ("initial", "verified_value", "exact_mismatches", "mismatches"), [(200.0, 4.0, 0, 0), (2.0, 4.0000005, 1, 0), (2.0, 5.0, 1, 1)], @@ -416,7 +412,6 @@ def test_sparse_delta_verification_compares_replacement( assert result["verification_mismatches"] == 2 * mismatches -@pytest.mark.vllm def test_fp8_weight_and_scale_use_exact_bit_overwrite() -> None: target = torch.tensor([0x38, 0x40, 0x48], dtype=torch.uint8).view( torch.float8_e4m3fn @@ -461,7 +456,6 @@ def test_fp8_weight_and_scale_use_exact_bit_overwrite() -> None: assert result["verification_abs_sum"] == 0.0 -@pytest.mark.vllm def test_xor_applies_exact_bits_and_replay_reverts() -> None: baseline = torch.tensor([1.0, 2.0, 3.0]) target = baseline.clone() @@ -483,7 +477,6 @@ def test_xor_applies_exact_bits_and_replay_reverts() -> None: assert torch.equal(_bits(target), _bits(baseline)) -@pytest.mark.vllm def test_overwrite_casts_absolute_source_values() -> None: target = torch.zeros(2, dtype=torch.float16) source = torch.tensor([1.25, -2.5], dtype=torch.float32) @@ -494,7 +487,6 @@ def test_overwrite_casts_absolute_source_values() -> None: assert torch.equal(target, source.to(torch.float16)) -@pytest.mark.vllm @pytest.mark.parametrize( ("name", "source", "targets", "error"), [ @@ -502,13 +494,13 @@ def test_overwrite_casts_absolute_source_values() -> None: "exp_transform", torch.tensor([math.log(2.0)]), {"exp_transform": torch.tensor([-1.0])}, - "transforms its input", + "without changing semantics", ), ( "weight", torch.tensor([1.0], dtype=torch.float32), {"identity": torch.zeros(1, dtype=torch.float16)}, - "dtypes must match", + "without changing semantics", ), ], ) @@ -524,7 +516,6 @@ def test_xor_rejects_non_bitwise_compatible_targets( _apply_payload(_applier(_NativeLoaderModel(**targets)), payload) -@pytest.mark.vllm def test_xor_rejects_overlapping_native_loader_copies() -> None: class RepeatedCopyModel(torch.nn.Module): def __init__(self) -> None: @@ -545,5 +536,5 @@ def load_weights(self, weights) -> None: "xor", ) - with pytest.raises(RuntimeError, match="overlapping target copies"): + with pytest.raises(RuntimeError, match="without changing semantics"): _apply_payload(_applier(RepeatedCopyModel()), payload) diff --git a/tests/unit/models/generation/test_vllm_sparse_refit.py b/tests/unit/models/generation/test_vllm_sparse_refit.py index 7f618817baa..bb766a221eb 100644 --- a/tests/unit/models/generation/test_vllm_sparse_refit.py +++ b/tests/unit/models/generation/test_vllm_sparse_refit.py @@ -31,24 +31,27 @@ VllmSparseRefitReceiver, _stage_sparse_payload, ) +from nemo_rl.models.generation.vllm.vllm_worker import VllmGenerationWorkerImpl from nemo_rl.models.generation.vllm.vllm_worker_async import ( VllmAsyncGenerationWorkerImpl, ) -from nemo_rl.utils.weight_transfer_remote_sparse import ( +from nemo_rl.utils.weight_transfer_http import ( G_VLLM_REFIT_API_KEY_HEADER, G_VLLM_REFIT_FLUSH_PATH, G_VLLM_REFIT_PREPARE_PATH, G_VLLM_REFIT_S3_MANIFEST_PATH, + G_VLLM_REFIT_ZMQ_FLUSH_PATH, ) from nemo_rl.utils.weight_transfer_sparse_codec import ( encode_sparse_infos, - iter_decoded_sparse_payload, + sparse_locations_for_item, ) from nemo_rl.utils.weight_transfer_zmq import ( G_VLLM_REFIT_CHECKSUM_HEADER, G_VLLM_REFIT_PAYLOAD_HEADER, G_VLLM_REFIT_PRODUCER_HEADER, G_VLLM_REFIT_TRANSFER_HEADER, + G_VLLM_REFIT_VERIFICATION_HEADER, G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, ) @@ -97,6 +100,15 @@ def _serialized_sparse_payload() -> bytes: return buffer.getvalue() +def _payload_locations(path: str) -> list[int]: + packed_locations, _, items = torch.load(path, weights_only=True, mmap=True) + return [ + int(location) + for item in items + for location in sparse_locations_for_item(item, packed_locations, device="cpu") + ] + + def _stage_payloads( receiver: VllmSparseRefitReceiver, staging_dir: Path, @@ -172,14 +184,19 @@ def test_sparse_refit_queue_deduplicates_transactional_payloads() -> None: receiver.update_weights_from_serialized_sparse_payloads = MagicMock( return_value={"ok": True, "payloads": 1} ) - assert receiver._enqueue_sparse_payload_apply(b"payload", key, "checksum")["ok"] - duplicate = receiver._enqueue_sparse_payload_apply(b"payload", key, "checksum") + assert receiver._enqueue_sparse_payload_apply(b"payload", key, "checksum", 4)[ + "ok" + ] + duplicate = receiver._enqueue_sparse_payload_apply( + b"payload", key, "checksum", 4 + ) assert duplicate == {"ok": True, "payloads": 0, "duplicate": True} with pytest.raises(ValueError, match="reused with different data"): receiver._enqueue_sparse_payload_apply(b"other", key, "different") response = receiver._flush_queued_sparse_payloads() assert response["payloads"] == 1 + assert response["verification_candidates"] == 4 assert receiver._refit_seen_payloads == {} @@ -255,7 +272,7 @@ def test_sparse_refit_queue_releases_condition_while_backpressured() -> None: assert pending_call.result(timeout=1.0)["ok"] -def test_sparse_refit_batch_decodes_once_before_collective_apply( +def test_sparse_refit_batch_stages_compact_payload_before_collective_apply( tmp_path: Path, ) -> None: with _sparse_refit_receiver() as receiver: @@ -264,13 +281,7 @@ def test_sparse_refit_batch_decodes_once_before_collective_apply( def collective_rpc(method: str, args: tuple[Any, ...]) -> list[Any]: assert method == "update_weights_from_decoded_sparse_payload" for path in args: - staged_locations.extend( - int(location) - for _, locations, _ in iter_decoded_sparse_payload( - torch.load(path, weights_only=True) - ) - for location in locations - ) + staged_locations.extend(_payload_locations(path)) return [ {"ok": True, "receiver_total_s": 1.0}, {"ok": True, "receiver_total_s": 1.0}, @@ -298,7 +309,7 @@ def collective_rpc(method: str, args: tuple[Any, ...]) -> list[Any]: assert response["payloads"] == 1 assert response["receiver_worker_total_s"] == 1.0 assert response["receiver_total_s"] >= 0.0 - assert receiver._refit_verification_candidates == 4 + assert receiver._refit_verification_candidates == 0 def test_sparse_refit_batch_drains_workers_before_error_cleanup(tmp_path: Path) -> None: @@ -368,13 +379,7 @@ async def collective_rpc( ) -> list[Any]: assert method == "update_weights_from_decoded_sparse_payload" for path in args: - staged_locations.extend( - location - for _, locations, _ in iter_decoded_sparse_payload( - torch.load(path, weights_only=True) - ) - for location in locations.tolist() - ) + staged_locations.extend(_payload_locations(path)) return [{"ok": True, "receiver_total_s": 1.0}] receiver._worker.llm = AsyncLlm() @@ -409,7 +414,11 @@ async def test_sparse_refit_payload_handlers_decode_and_enqueue(monkeypatch) -> ) s3_result = await receiver._apply_s3_manifest_payload( - {"key": "object-key", "checksum": "checksum"} + { + "key": "object-key", + "checksum": "checksum", + "verification_candidates": 4, + } ) assert s3_result["ok"] @@ -423,14 +432,15 @@ async def test_sparse_refit_payload_handlers_decode_and_enqueue(monkeypatch) -> G_VLLM_REFIT_PRODUCER_HEADER: "2", G_VLLM_REFIT_PAYLOAD_HEADER: "3", G_VLLM_REFIT_CHECKSUM_HEADER: "checksum", + G_VLLM_REFIT_VERIFICATION_HEADER: "5", }, body=AsyncMock(return_value=b"compressed"), ) zmq_result = await receiver._apply_zmq_payload(request) assert zmq_result["ok"] assert enqueue.call_args_list == [ - call(b"s3-payload", ("object-key", -1, -1), "checksum"), - call(b"zmq-payload", ("transfer", 2, 3), "checksum"), + call(b"s3-payload", ("object-key", -1, -1), "checksum", 4), + call(b"zmq-payload", ("transfer", 2, 3), "checksum", 5), ] @@ -455,6 +465,9 @@ def test_sparse_refit_api_auth_dispatch_and_error_mapping(monkeypatch) -> None: receiver._flush_queued_sparse_payloads = MagicMock( return_value={"ok": True, "payloads": 2} ) + receiver.flush_zmq_sparse_refit_relay = MagicMock( + return_value={"ok": True, "payloads": 3} + ) receiver._refit_collective_rpc = MagicMock(return_value=[]) app = FastAPI() receiver.setup_api_server(app) @@ -468,6 +481,11 @@ def test_sparse_refit_api_auth_dispatch_and_error_mapping(monkeypatch) -> None: headers=headers, ) flush_response = client.post(G_VLLM_REFIT_FLUSH_PATH, headers=headers) + zmq_flush_response = client.post( + G_VLLM_REFIT_ZMQ_FLUSH_PATH, + json={"transfer_id": "transfer"}, + headers=headers, + ) prepare_response = client.post( G_VLLM_REFIT_PREPARE_PATH, json={"tensors": {"weight": [[2, 3], "bfloat16"]}}, @@ -482,12 +500,14 @@ def test_sparse_refit_api_auth_dispatch_and_error_mapping(monkeypatch) -> None: assert unauthorized.status_code == 403 assert s3_response.status_code == 200 assert flush_response.status_code == 200 + assert zmq_flush_response.status_code == 200 assert prepare_response.status_code == 200 assert zmq_response.status_code == 500 assert receiver._refit_async_loop is not None receiver._apply_s3_manifest_payload.assert_awaited_once_with({"key": "key"}) receiver._apply_zmq_payload.assert_awaited_once() receiver._flush_queued_sparse_payloads.assert_called_once_with() + receiver.flush_zmq_sparse_refit_relay.assert_called_once_with("transfer", 0) receiver._refit_collective_rpc.assert_called_once_with( "prepare_sparse_delta_refit_info", ({"weight": ((2, 3), torch.bfloat16)},), @@ -590,11 +610,38 @@ def test_async_sparse_refit_exposes_zmq_relay(monkeypatch) -> None: } with _sparse_refit_receiver(async_engine=True, config=config) as receiver: + receiver._worker.base_url = "http://10.0.0.1:8000/v1" assert receiver.start_zmq_sparse_refit_relay(["http://receiver"]) == ( "tcp://10.0.0.1:12345" ) server.start.assert_called_once_with() + receiver.configure_zmq_sparse_refit_relay( + ["tcp://10.0.0.1:12345", "tcp://10.0.0.2:12345"] + ) + server.configure_tree.assert_called_once_with( + ["tcp://10.0.0.1:12345", "tcp://10.0.0.2:12345"], + own_address="tcp://10.0.0.1:12345", + local_refit_url="http://10.0.0.1:8000", + ) + server.flush.return_value = {"ok": True, "payloads": 2} + assert receiver.flush_zmq_sparse_refit_relay("transfer") == { + "ok": True, + "payloads": 2, + } + server.flush.assert_called_once_with("transfer", 0) receiver.stop_zmq_sparse_refit_relay() server.close.assert_called_once_with() assert receiver._zmq_refit_server is None + + +def test_vllm_worker_configures_zmq_relay() -> None: + worker = VllmGenerationWorkerImpl.__new__(VllmGenerationWorkerImpl) + worker._sparse_refit_receiver = MagicMock() + addresses = ["tcp://relay-0:19090", "tcp://relay-1:19090"] + + worker.configure_zmq_sparse_refit_relay(addresses) + + worker._sparse_refit_receiver.configure_zmq_sparse_refit_relay.assert_called_once_with( + addresses + ) diff --git a/tests/unit/models/policy/test_megatron_remote_sparse_refit.py b/tests/unit/models/policy/test_megatron_remote_sparse_refit.py index 4ea1dc34123..f3d5c62a7a4 100644 --- a/tests/unit/models/policy/test_megatron_remote_sparse_refit.py +++ b/tests/unit/models/policy/test_megatron_remote_sparse_refit.py @@ -16,100 +16,68 @@ import torch -from nemo_rl.models.policy.workers import megatron_remote_sparse_refit +from nemo_rl.models.generation.vllm.config import VllmDeltaCompressionConfig +from nemo_rl.models.policy.workers.megatron_remote_sparse_refit import ( + MegatronRemoteSparseRefit, +) -MegatronRemoteSparseRefit = megatron_remote_sparse_refit.MegatronRemoteSparseRefit +_DELTA_CONFIG = VllmDeltaCompressionConfig( + encoding="overwrite", sparse_bucket_size_bytes=1024 +) -_DELTA_CONFIG = {"encoding": "overwrite", "sparse_bucket_size_bytes": 1024} -_XOR_CONFIG = {**_DELTA_CONFIG, "encoding": "xor"} +def _worker(weights=()): + def export(): + return iter(weights) + return SimpleNamespace(_iter_params_with_optional_kv_scales=export) -class _AutoMapping: - is_expert = False - is_adapter = False - is_grouped_export = False - ep_size = 1 - ep_rank = 0 - tp_size = 1 - tp_rank = 0 - - def __init__(self, hf_param, *, parallelism="column", permute_dims=None): - self.hf_param = hf_param - self.parallelism = parallelism - self.permute_dims = permute_dims - - def _detect_parallelism_type(self, _module): - return self.parallelism - - -_ColumnMapping = type("_ColumnMapping", (_AutoMapping,), {}) -_DirectMapping = type("_DirectMapping", (_AutoMapping,), {}) - - -class _GatedMapping(_AutoMapping): - def __init__(self, *, gate, up): - super().__init__({"gate": gate, "up": up}) - - -_RowMapping = type("_RowMapping", (_AutoMapping,), {}) -_ReplicatedMapping = type("_ReplicatedMapping", (_AutoMapping,), {}) +def test_remote_sparse_initializes_canonical_hf_baseline(monkeypatch) -> None: + monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") + weights = [ + ("embedding.weight", torch.ones(2, 3)), + ("linear.weight", torch.ones(4, 3)), + ] + remote_refit = MegatronRemoteSparseRefit(_worker(weights), _DELTA_CONFIG) -def _install_mapping_types(monkeypatch, remote_refit_type): - monkeypatch.setattr( - remote_refit_type, - "_bridge_mapping_types", - staticmethod( - lambda: ( - _AutoMapping, - { - _ColumnMapping: megatron_remote_sparse_refit._COLUMN, - _DirectMapping: megatron_remote_sparse_refit._REPLICATED, - _GatedMapping: megatron_remote_sparse_refit._GATED, - _ReplicatedMapping: megatron_remote_sparse_refit._REPLICATED, - _RowMapping: megatron_remote_sparse_refit._ROW, - }, - ) - ), - ) - monkeypatch.setattr( - remote_refit_type, - "_bridge_exports_are_identity", - lambda _self: True, + info = remote_refit.initialize_baseline( + shard_rank=0, shard_count=1, transport="zmq" ) + assert info == { + name: (tuple(tensor.shape), tensor.dtype) for name, tensor in weights + } + assert set(vars(remote_refit)) == {"_worker", "_tracker"} -def _worker(tasks=(), *, fp8_cfg=None, export=None): - if export is None: - - def export(*, conversion_tasks=None): - return iter(()) - return SimpleNamespace( - cfg={}, - fp8_cfg=fp8_cfg, - model=object(), - megatron_bridge=SimpleNamespace( - get_conversion_tasks=lambda _models: list(tasks) - ), - _iter_params_with_optional_kv_scales=export, +def test_remote_sparse_preserves_xor_config() -> None: + remote_refit = MegatronRemoteSparseRefit( + _worker(), _DELTA_CONFIG.model_copy(update={"encoding": "xor"}) ) + assert remote_refit._tracker.encoding == "xor" -def test_remote_sparse_stream_drains_cuda_before_return(monkeypatch): - def export(*, conversion_tasks=None): - assert conversion_tasks == [] - return iter(()) - worker = _worker(export=export) - remote_refit = MegatronRemoteSparseRefit(worker, _DELTA_CONFIG) - result = {"payloads": 1, "changed_elements": 2, "total_elements": 3} +def test_remote_sparse_streams_one_canonical_path_and_drains_cuda(monkeypatch) -> None: + weights = [("model.weight", torch.ones(2))] + remote_refit = MegatronRemoteSparseRefit(_worker(weights), _DELTA_CONFIG) + expected = {"payloads": 1, "changed_elements": 2, "total_elements": 2} events = [] - def stream(*_args, **_kwargs): + def stream(iterator, **kwargs): + assert list(iterator) == weights + assert kwargs == { + "delta_tracker": remote_refit._tracker, + "transfer_id": "transfer", + "refit_targets": ["tcp://receiver:5555"], + "api_key_env_var": None, + "timeout_s": 1.0, + "shard_rank": 0, + "shard_count": 1, + } events.append("stream") - return result + return expected monkeypatch.setattr( "nemo_rl.models.policy.workers.megatron_remote_sparse_refit." @@ -118,10 +86,8 @@ def stream(*_args, **_kwargs): ) monkeypatch.setattr(torch.cuda, "is_available", lambda: True) monkeypatch.setattr(torch.cuda, "synchronize", lambda: events.append("sync")) - monkeypatch.setattr(torch.cuda.nvtx, "range_push", lambda _name: None) - monkeypatch.setattr(torch.cuda.nvtx, "range_pop", lambda: None) - actual = remote_refit.stream( + result = remote_refit.stream( "zmq", ["tcp://receiver:5555"], transfer_id="transfer", @@ -129,318 +95,29 @@ def stream(*_args, **_kwargs): timeout_s=1.0, shard_rank=0, shard_count=1, + overwrite_names=["model.weight"], ) - assert actual is result + assert result is expected + assert remote_refit._tracker.overwrite_names == frozenset({"model.weight"}) assert events == ["stream", "sync"] -def test_remote_sparse_stream_combines_local_and_misc_paths(monkeypatch): +def test_remote_sparse_finishes_single_tracker(monkeypatch) -> None: remote_refit = MegatronRemoteSparseRefit(_worker(), _DELTA_CONFIG) - remote_refit._policy_tracker = object() - remote_refit._local_tensors = [("local", torch.ones(1))] - remote_refit._misc_conversion_tasks = [] - monkeypatch.setattr(remote_refit, "_changed_misc_tasks", lambda: ([], 5, 6)) - calls = [] - - def stream(*_args, **kwargs): - calls.append(kwargs) - if kwargs["partition"] == "names": - return {"payloads": 4, "changed_elements": 5, "total_elements": 6} - return {"payloads": 1, "changed_elements": 2, "total_elements": 3} - - monkeypatch.setattr( - "nemo_rl.models.policy.workers.megatron_remote_sparse_refit." - "stream_sparse_delta_payloads_via_zmq", - stream, - ) - monkeypatch.setattr(torch.cuda, "is_available", lambda: False) - - result = remote_refit.stream( - "zmq", - ["tcp://receiver:5555"], - transfer_id="transfer", - api_key_env_var=None, - timeout_s=1.0, - shard_rank=2, - shard_count=4, - ) - - assert result == {"payloads": 5, "changed_elements": 7, "total_elements": 9} - assert {call["transfer_id"] for call in calls} == { - "transfer-local", - "transfer-misc", - } - local = next(call for call in calls if call["transfer_id"].endswith("-local")) - assert local["partition"] == "none" - - -def test_remote_sparse_globalizes_expert_name(): - task = SimpleNamespace( - mapping=SimpleNamespace(is_expert=True, ep_size=4, ep_rank=2), - megatron_module=SimpleNamespace(config=SimpleNamespace(num_moe_experts=16)), - ) - - assert ( - MegatronRemoteSparseRefit._canonical_hf_name( - task, "model.layers.0.mlp.experts.1.up_proj.weight" - ) - == "model.layers.0.mlp.experts.9.up_proj.weight" - ) - assert ( - MegatronRemoteSparseRefit._canonical_hf_name( - task, "model.layers.0.mlp.experts.9.up_proj.weight" - ) - == "model.layers.0.mlp.experts.9.up_proj.weight" - ) - - -def test_remote_sparse_projects_bridge_affine_mappings(monkeypatch): - _install_mapping_types(monkeypatch, MegatronRemoteSparseRefit) - cases = ( - ( - _AutoMapping("backbone.layers.0.mixer.D"), - torch.arange(8).view(4, 2), - "decoder.layers.0.mixer.D", - megatron_remote_sparse_refit._COLUMN, - ("backbone.layers.0.mixer.D", (8, 2), 0, 4), - ), - ( - _AutoMapping("backbone.layers.0.mixer.o_proj.weight", parallelism="row"), - torch.arange(8).view(2, 4), - "decoder.layers.0.self_attention.linear_proj.weight", - megatron_remote_sparse_refit._ROW, - ("backbone.layers.0.mixer.o_proj.weight", (2, 8), 1, 4), - ), - ( - _AutoMapping("backbone.layers.0.norm.weight", parallelism="replicated"), - torch.arange(4), - "decoder.layers.0.input_layernorm.weight", - megatron_remote_sparse_refit._REPLICATED, - ("backbone.layers.0.norm.weight", (4,), None, 0), - ), - ) - for mapping, tensor, global_name, kind, expected in cases: - mapping.tp_size = 2 - mapping.tp_rank = 1 - task = SimpleNamespace( - mapping=mapping, - megatron_module=torch.nn.Module(), - param_weight=tensor, - global_param_name=global_name, - ) - projection = MegatronRemoteSparseRefit._task_local_tensors(task, kind)[0][2] - assert ( - projection.name, - projection.global_shape, - projection.shard_dim, - projection.offset, - ) == expected - - gated = SimpleNamespace( - mapping=_GatedMapping( - gate="model.mlp.gate_proj.weight", - up="model.mlp.up_proj.weight", - ), - megatron_module=torch.nn.Module(), - param_weight=torch.arange(16).view(8, 2), - global_param_name="mlp.linear_fc1.weight", - ) - gated_projections = MegatronRemoteSparseRefit._task_local_tensors( - gated, megatron_remote_sparse_refit._GATED - ) - assert [projection.name for _, _, projection in gated_projections] == [ - "model.mlp.gate_proj.weight", - "model.mlp.up_proj.weight", - ] - assert [tuple(tensor.shape) for _, tensor, _ in gated_projections] == [ - (4, 2), - (4, 2), - ] - - -def test_remote_sparse_uses_local_baseline_to_gate_transformed_tasks( - monkeypatch, -): - class _TransformedMapping(_AutoMapping): - pass - - tensor = torch.tensor([1.0, 2.0, 3.0]) - task = SimpleNamespace( - mapping=_TransformedMapping("model.q_proj.weight"), - megatron_module=torch.nn.Linear(3, 1), - param_weight=tensor, - global_param_name="decoder.linear_qkv.weight", - ) - - _install_mapping_types(monkeypatch, MegatronRemoteSparseRefit) - monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") - remote_refit = MegatronRemoteSparseRefit(_worker([task]), _DELTA_CONFIG) - remote_refit._prepare_paths() - - assert remote_refit._misc_conversion_tasks == [task] - assert remote_refit._misc_local_tensors == [("0:decoder.linear_qkv.weight", tensor)] - assert remote_refit._policy_tracker is not None - remote_refit._policy_tracker.snapshot_baseline(remote_refit._misc_local_tensors) - - tensor[1] = 5 - changed_tasks, changed, total = remote_refit._changed_misc_tasks() - assert changed_tasks == [task] - assert (changed, total) == (1, 3) - - remote_refit._policy_tracker.on_sync_succeeded() - assert remote_refit._changed_misc_tasks() == ([], 0, 3) - - -def test_remote_sparse_uses_xor_only_for_direct_bitwise_path(monkeypatch): - class _TransformedMapping(_AutoMapping): - pass - - direct = torch.tensor([[1.0, 2.0], [3.0, 4.0]]) - transformed = torch.tensor([1.0, 2.0]) - tasks = [ - SimpleNamespace( - mapping=_AutoMapping("model.q_proj.weight"), - megatron_module=torch.nn.Linear(2, 2), - param_weight=direct, - global_param_name="decoder.linear_qkv.weight", - ), - SimpleNamespace( - mapping=_TransformedMapping("backbone.layers.0.mixer.A_log"), - megatron_module=torch.nn.Linear(2, 1), - param_weight=transformed, - global_param_name="decoder.mixer.A_log", - ), - ] - - _install_mapping_types(monkeypatch, MegatronRemoteSparseRefit) - monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") - remote_refit = MegatronRemoteSparseRefit(_worker(tasks), _XOR_CONFIG) - remote_refit._prepare_paths() - - assert remote_refit._policy_tracker is not None - assert remote_refit._policy_tracker.encoding == "xor" - assert remote_refit._tracker.encoding == "overwrite" - - remote_refit._policy_tracker.snapshot_baseline(remote_refit._local_tensors) - direct[0, 0] = 5 - direct_metadata = remote_refit._policy_tracker.prepare_sparse_delta_payload( - remote_refit._local_tensors - )[0][2] - assert [item["operation"] for item in direct_metadata] == ["xor"] - - residual = [("backbone.layers.0.mixer.A_log", transformed)] - remote_refit._tracker.snapshot_baseline(residual) - transformed[0] = 3 - residual_metadata = remote_refit._tracker.prepare_sparse_delta_payload(residual)[0][ - 2 - ] - assert [item["operation"] for item in residual_metadata] == ["overwrite"] - assert remote_refit.refit_info() == { - "backbone.layers.0.mixer.A_log": ((2,), torch.float32), - "model.q_proj.weight": ((2, 2), torch.float32), - } - - -def test_remote_sparse_preserves_bridge_task_dependencies(monkeypatch): - grouped = SimpleNamespace(is_grouped_export=True, group_key="experts") - tasks = [ - SimpleNamespace(mapping=grouped), - SimpleNamespace(mapping=grouped), - SimpleNamespace( - mapping=SimpleNamespace(is_grouped_export=False, group_key="other") - ), - ] - remote_refit = MegatronRemoteSparseRefit(object(), _DELTA_CONFIG) - remote_refit._misc_conversion_tasks = tasks - remote_refit._filter_misc_tasks = True - monkeypatch.setattr(remote_refit, "_all_reduce_max", lambda _flags: [1, 0, 0]) - - assert remote_refit._changed_misc_tasks()[0] == tasks[:2] - - remote_refit._filter_misc_tasks = False - assert remote_refit._changed_misc_tasks()[0] == tasks - - -def test_remote_sparse_balances_tasks_across_equivalent_replicas(monkeypatch): - from megatron.core import parallel_state - - from nemo_rl.utils.weight_transfer_remote_sparse import sparse_name_shard - - ranks = {"dp": 0, "expert_dp": 0} - monkeypatch.setattr(torch.distributed, "is_initialized", lambda: True) - monkeypatch.setattr( - parallel_state, - "get_data_parallel_rank", - lambda *, with_context_parallel: ranks["dp"], - ) - monkeypatch.setattr( - parallel_state, - "get_data_parallel_world_size", - lambda *, with_context_parallel: 2, - ) - monkeypatch.setattr( - parallel_state, - "get_expert_data_parallel_rank", - lambda: ranks["expert_dp"], - ) + events = [] monkeypatch.setattr( - parallel_state, - "get_expert_data_parallel_world_size", - lambda: 4, - ) - - dense = SimpleNamespace( - global_param_name="decoder.layers.0.input_layernorm.weight", - mapping=SimpleNamespace(is_expert=False, tp_rank=0, tp_size=2), + remote_refit._tracker, + "on_sync_succeeded", + lambda: events.append("succeeded"), ) - dense_owners = [] - for dp_rank in range(2): - ranks["dp"] = dp_rank - for tp_rank in range(2): - dense.mapping.tp_rank = tp_rank - if MegatronRemoteSparseRefit._owns_policy_local_task( - dense, replicated=True - ): - dense_owners.append(dp_rank * 2 + tp_rank) - assert dense_owners == [sparse_name_shard(dense.global_param_name, 4)] - - owner_counts = [0] * 4 - for expert_id in range(64): - name = f"decoder.layers.0.mlp.experts.local_experts.{expert_id}.weight" - expert = SimpleNamespace( - global_param_name=name, - mapping=SimpleNamespace(is_expert=True), - ) - owners = [] - for expert_dp_rank in range(4): - ranks["expert_dp"] = expert_dp_rank - if MegatronRemoteSparseRefit._owns_policy_local_task(expert): - owners.append(expert_dp_rank) - owner_counts[expert_dp_rank] += 1 - assert owners == [sparse_name_shard(name, 4)] - assert min(owner_counts) > 0 - - -def test_remote_sparse_fp8_policy_keeps_full_export_path(monkeypatch): - task = object() - - def export(*, conversion_tasks): - assert conversion_tasks == [task] - return iter(()) - - remote_refit = MegatronRemoteSparseRefit( - _worker([task], fp8_cfg={"fp8_param": True}, export=export), _DELTA_CONFIG - ) - snapshots = [] monkeypatch.setattr( - "nemo_rl.models.policy.workers.megatron_remote_sparse_refit." - "init_sparse_delta_baseline_from_iterator", - lambda iterator, **_kwargs: snapshots.append(list(iterator)), + remote_refit._tracker, + "on_sync_failed", + lambda: events.append("failed"), ) - remote_refit.initialize_baseline(shard_rank=0, shard_count=1, transport="zmq") + remote_refit.finish(True) + remote_refit.finish(False) - assert snapshots == [[]] - assert remote_refit._misc_conversion_tasks == [task] - assert remote_refit._policy_tracker is None + assert events == ["succeeded", "failed"] diff --git a/tests/unit/reference_configs/grpo_math_1B.yaml b/tests/unit/reference_configs/grpo_math_1B.yaml index 217556bfc2a..9dcabd2fc79 100644 --- a/tests/unit/reference_configs/grpo_math_1B.yaml +++ b/tests/unit/reference_configs/grpo_math_1B.yaml @@ -333,7 +333,7 @@ policy: stop_token_ids: null stop_strings: null refit_transport: null # Set to "vllm_s3_sparse" or "vllm_zmq_sparse" for remote sparse-delta refit. - delta_compression: null # Remote sparse-delta config; null uses the existing refit path. + delta_compression: null # Set {} for XOR and the 512 MiB bucket; null uses the existing refit path. mcore_generation_config: async_engine: false max_model_len: ${policy.max_total_sequence_length} diff --git a/tests/unit/utils/test_weight_transfer_remote_sparse.py b/tests/unit/utils/test_weight_transfer_stream.py similarity index 74% rename from tests/unit/utils/test_weight_transfer_remote_sparse.py rename to tests/unit/utils/test_weight_transfer_stream.py index 7853ff08028..35c68de85b1 100644 --- a/tests/unit/utils/test_weight_transfer_remote_sparse.py +++ b/tests/unit/utils/test_weight_transfer_stream.py @@ -15,29 +15,36 @@ import io import json import threading +from concurrent.futures import ThreadPoolExecutor from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from types import SimpleNamespace import pytest +import requests import torch import zstandard -from nemo_rl.utils import weight_transfer_remote_sparse, weight_transfer_zmq -from nemo_rl.utils.weight_transfer_remote_sparse import ( - download_s3_refit_payload, - sparse_payload_checksum, +from nemo_rl.utils import ( + weight_transfer_http, + weight_transfer_stream, + weight_transfer_zmq, ) from nemo_rl.utils.weight_transfer_sparse_codec import ( DeltaCompressionTracker, - SparseShardProjection, encode_sparse_infos, sparse_locations_for_item, ) +from nemo_rl.utils.weight_transfer_stream import ( + download_s3_refit_payload, + sparse_export_chunk_size, + sparse_payload_checksum, +) from nemo_rl.utils.weight_transfer_zmq import ( G_VLLM_REFIT_CHECKSUM_HEADER, G_VLLM_REFIT_PAYLOAD_HEADER, G_VLLM_REFIT_PRODUCER_HEADER, G_VLLM_REFIT_TRANSFER_HEADER, + G_VLLM_REFIT_VERIFICATION_HEADER, ZmqSparseRefitClient, ZmqSparseRefitServer, ) @@ -74,16 +81,33 @@ def snapshot_baseline(self, chunk) -> None: self.names.extend(name for name, _tensor in chunk) +class _SparseTestTransport: + name = "zmq" + transfer_workers = 1 + + def __init__(self, send_payload) -> None: + self._send_payload = send_payload + self.cleaned = False + + def send(self, body, payload_id, _verification_candidates): + return self._send_payload(body, payload_id) + + def cleanup(self) -> None: + self.cleaned = True + + def _stream_sparse_test_payloads(tensors, send_payload): - return weight_transfer_remote_sparse.stream_sparse_delta_payloads( - tensors, - delta_tracker=_SparsePipelineTracker(), - transport="zmq", - send_payload=send_payload, - transfer_workers=1, - shard_rank=0, - shard_count=1, - ) + transport = _SparseTestTransport(send_payload) + try: + return weight_transfer_stream.stream_sparse_delta_payloads( + tensors, + delta_tracker=_SparsePipelineTracker(), + transport=transport, + shard_rank=0, + shard_count=1, + ) + finally: + assert transport.cleaned def _delta_tracker(encoding: str = "overwrite") -> DeltaCompressionTracker: @@ -92,15 +116,14 @@ def _delta_tracker(encoding: str = "overwrite") -> DeltaCompressionTracker: ) -def _baseline_names(tensors, *, rank: int, partition: str = "chunks"): +def _baseline_names(tensors, *, rank: int): tracker = _BaselineNamesTracker() - weight_transfer_remote_sparse.init_sparse_delta_baseline_from_iterator( + weight_transfer_stream.init_sparse_delta_baseline_from_iterator( tensors, delta_tracker=tracker, shard_rank=rank, shard_count=2, transport="zmq", - partition=partition, ) return tracker.names @@ -123,21 +146,6 @@ def test_delta_tracker_commits_only_successful_syncs() -> None: assert not tracker.prepare_sparse_delta_payload([("weight", tensor)])[0][2] -def test_delta_tracker_change_summary_is_transactional() -> None: - tracker = _delta_tracker() - first = torch.tensor([1.0, 2.0]) - second = torch.tensor([3.0, 4.0]) - tensors = [("first", first), ("second", second)] - tracker.snapshot_baseline(tensors) - second[0] = 5 - - assert tracker.prepare_change_summary(tensors) == ({"second"}, 1, 4) - tracker.on_sync_failed() - assert tracker.prepare_change_summary(tensors) == ({"second"}, 1, 4) - tracker.on_sync_succeeded() - assert tracker.prepare_change_summary(tensors) == (set(), 0, 4) - - def test_delta_tracker_emits_bounded_verification_budget(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "2") tracker = _delta_tracker() @@ -193,42 +201,20 @@ def test_delta_tracker_xor_encodes_against_baseline() -> None: assert torch.equal(tracker.baseline["weight"], tensor) -@pytest.mark.parametrize( - ("projection", "expected_locations"), - [ - (SparseShardProjection("hf.weight", (2, 4), 1, 2), [2, 7]), - (SparseShardProjection("hf.weight", (4, 2), 0, 2), [4, 7]), - ], -) -def test_delta_tracker_projects_local_shards_to_hf_locations( - monkeypatch, - projection: SparseShardProjection, - expected_locations: list[int], -) -> None: - monkeypatch.setenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "2") - tracker = DeltaCompressionTracker( - {"encoding": "overwrite", "sparse_bucket_size_bytes": 1024}, - projections={"local.weight": projection}, - ) - tensor = torch.zeros(2, 2) - tracker.snapshot_baseline([("local.weight", tensor)]) - tensor[0, 0] = 1 - tensor[1, 1] = 2 - - (locations, _, metadata), changed, total = tracker.prepare_sparse_delta_payload( - [("local.weight", tensor)] +def test_delta_tracker_uses_overwrite_for_receiver_incompatible_weights() -> None: + tracker = _delta_tracker("xor") + weight = torch.tensor([1.0]) + scale = torch.tensor([2.0]) + tracker.snapshot_baseline([("weight", weight), ("scale", scale)]) + tracker.overwrite_names = frozenset({"scale"}) + weight.add_(1) + scale.add_(1) + + (_, _, metadata), _, _ = tracker.prepare_sparse_delta_payload( + [("weight", weight), ("scale", scale)] ) - assert (changed, total) == (2, 4) - assert metadata[0]["name"] == "hf.weight" - assert metadata[0]["shape"] == projection.global_shape - assert metadata[0]["verification_samples"] == 2 - assert ( - sparse_locations_for_item(metadata[0], locations, device="cpu").tolist() - == expected_locations - ) - tracker.on_sync_succeeded() - assert not tracker.prepare_sparse_delta_payload([("local.weight", tensor)])[0][2] + assert [item["operation"] for item in metadata] == ["xor", "overwrite"] def test_sparse_index_encoding_preserves_uint64_locations() -> None: @@ -296,7 +282,7 @@ def test_delta_tracker_rejects_arithmetic_encoding() -> None: def test_s3_download_verifies_checksum(monkeypatch) -> None: compressed = zstandard.ZstdCompressor().compress(b"payload") monkeypatch.setattr( - weight_transfer_remote_sparse, + weight_transfer_stream, "_get_manifest_s3_store", lambda *_args: SimpleNamespace(get=lambda _key: bytearray(compressed)), ) @@ -314,16 +300,38 @@ def test_s3_download_verifies_checksum(monkeypatch) -> None: def test_refit_http_session_does_not_retry_application_errors() -> None: - retry = ( - weight_transfer_remote_sparse.refit_http_session() - .get_adapter("http://") - .max_retries - ) + retry = weight_transfer_http.refit_http_session().get_adapter("http://").max_retries assert 500 not in retry.status_forcelist assert {502, 503, 504} <= set(retry.status_forcelist) +def test_refit_http_sessions_share_connection_pool_across_threads() -> None: + barrier = threading.Barrier(4) + + def adapter(_): + barrier.wait() + return weight_transfer_http.refit_http_session().get_adapter("http://") + + with ThreadPoolExecutor(max_workers=4) as executor: + adapters = list(executor.map(adapter, range(4))) + + assert all(adapter is adapters[0] for adapter in adapters) + + +def test_refit_http_error_preserves_non_json_status_and_body(monkeypatch) -> None: + response = requests.Response() + response.status_code = 500 + response._content = b"gateway failure" + session = SimpleNamespace(post=lambda *_args, **_kwargs: response) + monkeypatch.setattr(weight_transfer_http, "refit_http_session", lambda: session) + + with pytest.raises(RuntimeError, match="HTTP 500: gateway failure"): + weight_transfer_http.post_vllm_refit_endpoints( + ["http://receiver/refit"], {}, api_key=None, timeout_s=1.0 + ) + + def test_sparse_export_finishes_before_blocked_transfers(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_ZMQ_ENCODE_WORKERS", "1") monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "1") @@ -374,6 +382,40 @@ def fail_transfer(_body, _payload_index): assert exported == list(range(4)) +def test_sparse_transport_cleanup_runs_on_transfer_workers(monkeypatch) -> None: + monkeypatch.setenv("NRL_REFIT_ZMQ_ENCODE_WORKERS", "1") + monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "1") + + class Transport: + name = "zmq" + transfer_workers = 2 + + def __init__(self) -> None: + self.barrier = threading.Barrier(2) + self.send_threads = set() + self.cleanup_threads = set() + + def send(self, _body, _payload_id, _verification_candidates): + self.send_threads.add(threading.get_ident()) + self.barrier.wait(timeout=5.0) + return {"receiver": {}} + + def cleanup(self) -> None: + self.cleanup_threads.add(threading.get_ident()) + + transport = Transport() + result = weight_transfer_stream.stream_sparse_delta_payloads( + [(f"weight-{index}", torch.ones(1)) for index in range(2)], + delta_tracker=_SparsePipelineTracker(), + transport=transport, + shard_rank=0, + shard_count=1, + ) + + assert result["payloads"] == 2 + assert transport.send_threads == transport.cleanup_threads + + def test_sparse_stream_coalesces_export_chunks(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_ZMQ_ENCODE_WORKERS", "2") monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "4") @@ -392,17 +434,16 @@ def send(body, payload_index): ) return {"receiver": {}} - result = weight_transfer_remote_sparse.stream_sparse_delta_payloads( + transport = _SparseTestTransport(send) + result = weight_transfer_stream.stream_sparse_delta_payloads( tensors, delta_tracker=tracker, - transport="zmq", - send_payload=send, - transfer_workers=1, + transport=transport, shard_rank=0, shard_count=1, - partition="none", ) + assert transport.cleaned assert result == {"payloads": 2, "changed_elements": 4, "total_elements": 4} payload_names = [ [item["name"] for item in payloads[index][2]] for index in sorted(payloads) @@ -421,6 +462,16 @@ def send(body, payload_index): ].tolist() == [1065353216] +def test_sparse_export_chunk_defaults_are_transport_specific(monkeypatch) -> None: + monkeypatch.delenv("NRL_REFIT_S3_EXPORT_CHUNK_BYTES", raising=False) + monkeypatch.delenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", raising=False) + tracker = _delta_tracker() + tracker.sparse_bucket_size_bytes = 1024**3 + + assert sparse_export_chunk_size(tracker, "s3") == 64 * 1024**2 + assert sparse_export_chunk_size(tracker, "zmq") == 256 * 1024**2 + + def test_sparse_baseline_snapshots_only_owned_export_chunks( monkeypatch, capsys ) -> None: @@ -429,27 +480,25 @@ def test_sparse_baseline_snapshots_only_owned_export_chunks( assert _baseline_names(tensors, rank=1) == ["weight-1", "weight-3"] assert "chunks=4" in capsys.readouterr().out - assert _baseline_names(tensors, rank=1, partition="none") == [ - f"weight-{index}" for index in range(4) - ] - -def test_sparse_name_partition_is_stable_for_filtered_exports(monkeypatch) -> None: - monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "4") - tensors = [(f"weight-{index}", torch.tensor([float(index)])) for index in range(8)] - owners = [ - set(_baseline_names(tensors, rank=rank, partition="names")) for rank in range(2) - ] +def test_sparse_stream_sends_only_owned_export_chunks(monkeypatch) -> None: + monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "1") + sent = [] + transport = _SparseTestTransport( + lambda _body, payload_id: sent.append(payload_id) or {"receiver": {}} + ) - assert owners[0].isdisjoint(owners[1]) - assert owners[0] | owners[1] == {name for name, _tensor in tensors} + result = weight_transfer_stream.stream_sparse_delta_payloads( + [(f"weight-{index}", torch.ones(1)) for index in range(4)], + delta_tracker=_SparsePipelineTracker(), + transport=transport, + shard_rank=1, + shard_count=2, + ) - filtered = tensors[::2] - for rank in range(2): - assert set(_baseline_names(filtered, rank=rank, partition="names")) == owners[ - rank - ] & {name for name, _tensor in filtered} + assert result == {"payloads": 2, "changed_elements": 2, "total_elements": 2} + assert sorted(sent) == [0, 1] def test_s3_manifest_transport_validates_configuration(monkeypatch) -> None: @@ -464,13 +513,13 @@ def test_s3_manifest_transport_validates_configuration(monkeypatch) -> None: } monkeypatch.setenv("NRL_REFIT_S3_BUCKET", "bucket") with pytest.raises(ValueError, match="URL is required"): - weight_transfer_remote_sparse.stream_sparse_delta_payloads_via_s3_manifest( + weight_transfer_stream.stream_sparse_delta_payloads_via_s3_manifest( refit_targets=[], **kwargs ) monkeypatch.delenv("NRL_REFIT_S3_BUCKET") with pytest.raises(RuntimeError, match="NRL_REFIT_S3_BUCKET"): - weight_transfer_remote_sparse.stream_sparse_delta_payloads_via_s3_manifest( + weight_transfer_stream.stream_sparse_delta_payloads_via_s3_manifest( refit_targets=["http://receiver"], **kwargs ) @@ -496,7 +545,7 @@ def delete(self, key) -> None: monkeypatch.setenv("NRL_REFIT_S3_UPLOAD_WORKERS", "3") monkeypatch.setenv("NRL_TEST_REFIT_KEY", "secret") monkeypatch.setattr( - weight_transfer_remote_sparse, + weight_transfer_stream, "_get_manifest_s3_store", lambda *_args: store, ) @@ -510,20 +559,18 @@ def post(endpoints, manifest, **kwargs): def stream(iterator, **kwargs): assert list(iterator) == [("weight", torch.tensor([1.0]))] - assert kwargs["transport"] == "s3" - assert kwargs["transfer_workers"] == 3 - response = kwargs["send_payload"](b"payload", 3) + transport = kwargs["transport"] + assert transport.name == "s3" + assert transport.transfer_workers == 3 + response = transport.send(b"payload", 3, 4) assert response["receiver"] == {"receiver_total_s": 2.0} + transport.cleanup() return {"payloads": 1, "changed_elements": 1, "total_elements": 1} - monkeypatch.setattr( - weight_transfer_remote_sparse, "post_vllm_refit_endpoints", post - ) - monkeypatch.setattr( - weight_transfer_remote_sparse, "stream_sparse_delta_payloads", stream - ) + monkeypatch.setattr(weight_transfer_stream, "post_vllm_refit_endpoints", post) + monkeypatch.setattr(weight_transfer_stream, "stream_sparse_delta_payloads", stream) - result = weight_transfer_remote_sparse.stream_sparse_delta_payloads_via_s3_manifest( + result = weight_transfer_stream.stream_sparse_delta_payloads_via_s3_manifest( [("weight", torch.tensor([1.0]))], delta_tracker=SimpleNamespace(), refit_targets=[" http://receiver-a/ ", "http://receiver-b"], @@ -548,18 +595,19 @@ def stream(iterator, **kwargs): "region": store.region, "key": key, "checksum": sparse_payload_checksum(b"payload"), + "verification_candidates": 4, }, {"api_key": "secret", "timeout_s": 7.0}, ) ] -def test_zmq_stream_routes_shards_and_reuses_clients(monkeypatch) -> None: +def test_zmq_stream_routes_shards_and_closes_clients(monkeypatch) -> None: created = [] sent = [] + closed = [] monkeypatch.setenv("NRL_TEST_REFIT_KEY", "secret") monkeypatch.setenv("NRL_REFIT_ZMQ_SEND_WORKERS", "2") - monkeypatch.delattr(weight_transfer_zmq._ZMQ_LOCAL, "clients", raising=False) class Client: def __init__(self, address, **kwargs) -> None: @@ -569,12 +617,19 @@ def send_payload(self, **kwargs): sent.append(kwargs) return {"ok": True, "receiver_total_s": 0.5} + def close(self) -> None: + closed.append(True) + def stream(_iterator, **kwargs): - assert kwargs["transport"] == "zmq" - assert kwargs["transfer_workers"] == 2 + transport = kwargs["transport"] + assert transport.name == "zmq" + assert transport.transfer_workers == 2 for payload_id in range(2): - response = kwargs["send_payload"](f"body-{payload_id}".encode(), payload_id) + response = transport.send( + f"body-{payload_id}".encode(), payload_id, payload_id + 1 + ) assert response["receiver"]["ok"] + transport.cleanup() return {"payloads": 2, "changed_elements": 2, "total_elements": 2} monkeypatch.setattr(weight_transfer_zmq, "ZmqSparseRefitClient", Client) @@ -611,13 +666,15 @@ def stream(_iterator, **kwargs): == 2 ) - assert created == [ + assert created == 2 * [ ( "tcp://receiver-b", {"timeout_s": 7.0, "producer_id": 3, "api_key": "secret"}, ) ] + assert closed == [True, True] assert [item["payload_id"] for item in sent] == [0, 1, 0, 1] + assert [item["verification_candidates"] for item in sent] == [1, 2, 1, 2] assert all( item["checksum"] == sparse_payload_checksum(item["body"]) for item in sent ) @@ -690,6 +747,7 @@ def test_zmq_server_rejects_malformed_messages() -> None: "producer_id": 0, "payload_id": 1, "checksum": sparse_payload_checksum(body), + "verification_candidates": 2, } def frames(kind: bytes = b"DATA", **updates: object) -> list[bytes]: @@ -741,34 +799,41 @@ def _send_zmq_payload( payload_id: int, body: bytes, checksum: str | None = None, + transfer_id: str = "transfer-a", ) -> dict[str, object]: return client.send_payload( - transfer_id="transfer-a", + transfer_id=transfer_id, payload_id=payload_id, checksum=checksum or sparse_payload_checksum(body), + verification_candidates=2, body=body, ) def test_zmq_sparse_refit_relay_fans_out_and_rejects_corruption(monkeypatch) -> None: monkeypatch.setenv("NRL_TEST_REFIT_KEY", "secret") - received = [[], []] + received = [[] for _ in range(4)] receivers = [_receiver_server(items) for items in received] urls = [f"http://127.0.0.1:{server.server_port}" for server, _ in receivers] - relay = ZmqSparseRefitServer( - urls, - bind_address="tcp://127.0.0.1:*", - api_key_env_var="NRL_TEST_REFIT_KEY", - timeout_s=5.0, - ) - address = relay.start() + relays = [ + ZmqSparseRefitServer( + urls, + bind_address="tcp://127.0.0.1:*", + api_key_env_var="NRL_TEST_REFIT_KEY", + timeout_s=5.0, + ) + for _ in urls + ] + addresses = [relay.start() for relay in relays] + for relay, address, url in zip(relays, addresses, urls, strict=True): + relay.configure_tree(addresses, own_address=address, local_refit_url=url) unauthenticated_client = ZmqSparseRefitClient( - address, + addresses[0], timeout_s=5.0, producer_id=2, ) client = ZmqSparseRefitClient( - address, + addresses[0], timeout_s=5.0, producer_id=3, api_key="secret", @@ -779,9 +844,14 @@ def test_zmq_sparse_refit_relay_fans_out_and_rejects_corruption(monkeypatch) -> with pytest.raises(RuntimeError, match="authentication failed"): _send_zmq_payload(unauthenticated_client, 6, body) first = _send_zmq_payload(client, 7, body) - assert first["ok"] - assert first["receiver_total_s"] == 0.25 - assert [len(items) for items in received] == [1, 1] + duplicate = _send_zmq_payload(client, 7, body) + assert first["ok"] and first["staged"] + assert duplicate["ok"] and duplicate["staged"] + for relay in relays: + flushed = relay.flush("transfer-a", expected_payloads=1) + assert flushed["payloads"] == 1 + assert flushed["receiver_total_s"] == 0.25 + assert [len(items) for items in received] == [1] * 4 for items in received: headers, posted_body = items[0] assert posted_body == body @@ -789,14 +859,19 @@ def test_zmq_sparse_refit_relay_fans_out_and_rejects_corruption(monkeypatch) -> assert headers[G_VLLM_REFIT_PRODUCER_HEADER] == "3" assert headers[G_VLLM_REFIT_PAYLOAD_HEADER] == "7" assert headers[G_VLLM_REFIT_CHECKSUM_HEADER] == checksum + assert headers[G_VLLM_REFIT_VERIFICATION_HEADER] == "2" assert headers["x-nemo-rl-refit-key"] == "secret" + with pytest.raises(RuntimeError, match="already flushed"): + _send_zmq_payload(client, 8, body) + assert _send_zmq_payload(client, 8, body, "0" * 32, "transfer-b")["staged"] with pytest.raises(RuntimeError, match="checksum mismatch"): - _send_zmq_payload(client, 8, body, "0" * 32) + relays[0].flush("transfer-b") finally: unauthenticated_client.close() client.close() - relay.close() + for relay in relays: + relay.close() for server, thread in receivers: server.shutdown() thread.join(timeout=5.0) 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 ebc55cbb4f1..c54c13696d3 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 @@ -12,10 +12,11 @@ # See the License for the specific language governing permissions and # limitations under the License. -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest +from nemo_rl.models.generation.vllm.config import VllmDeltaCompressionConfig from nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer import ( VllmRemoteSparseWeightSynchronizer, validate_vllm_remote_sparse_refit, @@ -51,15 +52,12 @@ def _remote_sparse_sync( policy.worker_group.run_all_workers_single_data.return_value = commit_refs generation = MagicMock() - generation.worker_group.run_all_workers_single_data.side_effect = [ - [MagicMock()], - *([[MagicMock()]] if transport == "zmq" else []), - ] + generation.worker_group.run_all_workers_single_data.return_value = [MagicMock()] generation.invalidate_kv_cache.return_value = True get_results: list[object] = [["http://receiver"]] if transport == "zmq": - get_results.append(["tcp://relay:19090"]) + get_results.extend([["tcp://relay:19090"], [None]]) get_results.extend([[{"weight": ((8,), "float32")}], stream_result]) mock_ray.get.side_effect = get_results @@ -78,12 +76,16 @@ def _valid_config() -> dict: def test_validate_remote_sparse_refit_accepts_supported_scope(): + config = _valid_config() assert ( validate_vllm_remote_sparse_refit( - _valid_config(), colocated=False, megatron_enabled=True + config, colocated=False, megatron_enabled=True ) == "s3" ) + assert config["delta_compression"] == VllmDeltaCompressionConfig( + encoding="overwrite" + ) @pytest.mark.parametrize( @@ -127,6 +129,7 @@ def test_init_communicator_joins_prelaunched_baseline(self, mock_ray, post): transport="s3", baseline_init_refs=[baseline_ref], ) + post.return_value = [{"overwrite_names": ["weight"]}] sync.init_communicator() @@ -139,6 +142,7 @@ def test_init_communicator_joins_prelaunched_baseline(self, mock_ray, post): timeout_s=600.0, ) assert sync._baseline_init_refs == [] + assert sync._overwrite_names == ["weight"] def test_merge_refit_info_rejects_conflicting_metadata(self): with pytest.raises(ValueError, match="Conflicting sparse refit metadata"): @@ -190,6 +194,7 @@ def test_shutdown_cancels_pending_work_and_stops_zmq(self, mock_ray): assert sync._baseline_commit_refs == [] assert sync._refit_urls == [] assert sync._targets == [] + assert sync._overwrite_names == [] def test_fails_before_transfer_when_kv_cache_invalidation_fails(self, mock_ray): policy = MagicMock() @@ -219,25 +224,48 @@ def test_initializes_streams_commits_and_updates_baseline( "verification_max_abs": 1e-9, } ] + verification = post.return_value + post.side_effect = [ + [{"ok": True, "payloads": 3, "receiver_relay_fanout_s": 1.0}], + verification, + ] metrics = sync.sync_weights() assert [ entry.args[0] for entry in policy.worker_group.run_all_workers_multiple_data.call_args_list ] == ["init_remote_sparse_delta_baseline", "stream_remote_sparse_weights"] + assert ( + policy.worker_group.run_all_workers_multiple_data.call_args_list[1].kwargs[ + "common_kwargs" + ]["overwrite_names"] + == [] + ) assert [ entry.args[0] for entry in generation.worker_group.run_all_workers_single_data.call_args_list ] == [ "report_refit_server_base_url", "start_zmq_sparse_refit_relay", + "configure_zmq_sparse_refit_relay", ] generation.worker_group.run_all_workers_single_data.assert_any_call( "start_zmq_sparse_refit_relay", run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], refit_urls=["http://receiver"], ) - post.assert_called_once_with( + generation.worker_group.run_all_workers_single_data.assert_any_call( + "configure_zmq_sparse_refit_relay", + run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], + relay_addresses=["tcp://relay:19090"], + ) + post.assert_any_call( + ["http://receiver/nemo-rl/refit/zmq-flush"], + {"transfer_id": ANY, "expected_payloads": 3}, + api_key=None, + timeout_s=600.0, + ) + post.assert_any_call( ["http://receiver/nemo-rl/refit/flush"], {}, api_key=None, @@ -258,6 +286,7 @@ def test_initializes_streams_commits_and_updates_baseline( assert metrics["delta_verify/mean_abs"] == 2.5e-10 assert metrics["delta_verify/max_abs"] == 1e-9 assert metrics["transfer/payloads"] == 3.0 + assert metrics["transfer/relay_flush_s"] >= 0.0 assert not sync.is_stale def test_sample_mismatch_does_not_commit_baseline(self, mock_ray, post): @@ -275,6 +304,12 @@ def test_sample_mismatch_does_not_commit_baseline(self, mock_ray, post): } ] + verification = post.return_value + post.side_effect = [ + [{"ok": True, "payloads": 3}], + verification, + verification, + ] with pytest.raises(RuntimeError, match="1 mismatched deltas out of 4"): sync.sync_weights() @@ -299,3 +334,30 @@ def test_failure_drains_receivers_without_committing_baseline(self, mock_ray, po policy.worker_group.run_all_workers_single_data.assert_called_once_with( "finish_remote_sparse_delta_sync", succeeded=False ) + + def test_zmq_relay_failure_drains_before_rejecting_baseline(self, mock_ray, post): + sync, policy, _ = _remote_sparse_sync( + mock_ray, + "zmq", + [{"payloads": 3, "changed_elements": 3, "total_elements": 100}], + ) + mock_ray.get.side_effect = [ + [{"payloads": 3, "changed_elements": 3, "total_elements": 100}], + ] + post.side_effect = [ + RuntimeError("fanout failed"), + [{"ok": True, "payloads": 3}], + [{"ok": True}], + ] + + with pytest.raises(RuntimeError, match="fanout failed"): + sync.sync_weights() + + assert [call.args[0][0] for call in post.call_args_list] == [ + "http://receiver/nemo-rl/refit/zmq-flush", + "http://receiver/nemo-rl/refit/zmq-flush", + "http://receiver/nemo-rl/refit/flush", + ] + policy.worker_group.run_all_workers_single_data.assert_called_once_with( + "finish_remote_sparse_delta_sync", succeeded=False + ) diff --git a/tools/refit_bandwidth_calculator.py b/tools/refit_bandwidth_calculator.py index 2dbeb34fbc9..1a2a69f7258 100644 --- a/tools/refit_bandwidth_calculator.py +++ b/tools/refit_bandwidth_calculator.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Compare measured zstd sparse refit with NCCL projected to Ethernet.""" +"""Project current zstd sparse refit against NCCL over candidate Ethernet.""" import argparse import json @@ -25,16 +25,16 @@ _REFERENCE_IB_GBPS = 400.0 _DENSITIES = (3.0, 5.0) -_SPARSE_SIZE_RANGE_GB = (63.2, 1121.0) -# Each pair is (fixed seconds, seconds per 1,000 GB) at 3% and 5% changed. -_SPARSE_LATENCY_FITS: dict[ - Transport, tuple[tuple[float, float], tuple[float, float]] -] = { - "s3": ((2.271898, 90.539029), (9.132526, 164.284601)), - "zmq": ((5.290113, 78.016171), (10.586558, 124.416925)), +_SPARSE_ANCHOR_SIZE_GB = 247.2 +_SPARSE_BUCKET_SIZE_BYTES = 512 * 1024**2 +_SPARSE_ANCHOR_LATENCY_S: dict[Transport, tuple[float, float]] = { + "s3": (20.233733, 25.790387), + "zmq": (24.095243, 33.7759615), +} +_SPARSE_ANCHOR_WIRE_GB: dict[Transport, tuple[float, float]] = { + "s3": (5.858410822107136, 9.765203805732864), + "zmq": (5.608205511, 9.3456392505), } -# Unique compressed wire GB per 1,000 GB of indexed BF16 weights at 3% and 5%. -_SPARSE_WIRE_GB_PER_TB = (22.430807, 37.362273) _NCCL_ANCHORS = ( (63.2, 0.84, 1.60), (247.2, 1.46, 1.74), @@ -48,6 +48,7 @@ class Estimate: transport: Transport model_size_gb: float changed_pct: float + sparse_bucket_size_bytes: int sparse_seconds: float approximate_wire_gb: float nccl_ib_low_s: float @@ -60,29 +61,35 @@ class Estimate: candidate_winner: str | None +def _project_density(anchors: tuple[float, float], changed_pct: float) -> float: + exponent = math.log(anchors[1] / anchors[0]) / math.log(5.0 / 3.0) + return anchors[0] * (changed_pct / 3.0) ** exponent + + def predict_sparse_seconds( model_size_gb: float, changed_pct: float, *, transport: Transport, ) -> float: - """Evaluate the campaign fit at any positive changed density.""" + """Linearly project the latest 120B measurement by model size.""" if model_size_gb <= 0 or changed_pct <= 0: raise ValueError("model_size_gb and changed_pct must be positive") - size_tb = model_size_gb / 1000.0 - (a3, b3), (a5, b5) = _SPARSE_LATENCY_FITS[transport] - latency_3, latency_5 = a3 + b3 * size_tb, a5 + b5 * size_tb - exponent = math.log(latency_5 / latency_3) / math.log(5.0 / 3.0) - return latency_3 * (changed_pct / 3.0) ** exponent + anchor = _project_density(_SPARSE_ANCHOR_LATENCY_S[transport], changed_pct) + return anchor * model_size_gb / _SPARSE_ANCHOR_SIZE_GB -def predict_sparse_wire_gb(model_size_gb: float, changed_pct: float) -> float: - """Scale the measured zstd wire fit to model size and changed density.""" +def predict_sparse_wire_gb( + model_size_gb: float, + changed_pct: float, + *, + transport: Transport, +) -> float: + """Scale the latest measured zstd wire bytes by model size and density.""" if model_size_gb <= 0 or changed_pct <= 0: raise ValueError("model_size_gb and changed_pct must be positive") - wire_3, wire_5 = _SPARSE_WIRE_GB_PER_TB - exponent = math.log(wire_5 / wire_3) / math.log(5.0 / 3.0) - return model_size_gb / 1000.0 * wire_3 * (changed_pct / 3.0) ** exponent + anchor = _project_density(_SPARSE_ANCHOR_WIRE_GB[transport], changed_pct) + return anchor * model_size_gb / _SPARSE_ANCHOR_SIZE_GB def _nccl_reference(model_size_gb: float) -> tuple[float, float]: @@ -134,8 +141,13 @@ def estimate( transport, model_size_gb, changed_pct, + _SPARSE_BUCKET_SIZE_BYTES, sparse_seconds, - predict_sparse_wire_gb(model_size_gb, changed_pct), + predict_sparse_wire_gb( + model_size_gb, + changed_pct, + transport=transport, + ), nccl_low, nccl_high, candidate_ethernet_gbps, @@ -162,7 +174,8 @@ def _print_results(results: list[Estimate]) -> None: first = results[0] print( f"Model: {first.model_size_gb:g} GB indexed BF16; " - f"changed: {first.changed_pct:g}%; compression: zstd" + f"changed: {first.changed_pct:g}%; compression: zstd; " + f"bucket: {first.sparse_bucket_size_bytes // 1024**2} MiB" ) print( "Measured NCCL on 400 Gbps/rank H100 IB: " @@ -193,8 +206,11 @@ def _print_results(results: list[Estimate]) -> None: "\nBelow the lower crossover sparse wins across the NCCL envelope; " "above the upper crossover NCCL wins." ) - if not _SPARSE_SIZE_RANGE_GB[0] <= first.model_size_gb <= _SPARSE_SIZE_RANGE_GB[1]: - print("Note: model size is outside the measured sparse calibration range.") + if not math.isclose(first.model_size_gb, _SPARSE_ANCHOR_SIZE_GB): + print( + f"Note: sparse time is a linear model-size projection from the " + f"{_SPARSE_ANCHOR_SIZE_GB:g} GB measured anchor." + ) if not _DENSITIES[0] <= first.changed_pct <= _DENSITIES[1]: print("Note: changed density is extrapolated from measured 3% and 5% arms.") From f7c1eec0563c1b73c2e0bcbb84ba0c091c94fd04 Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Thu, 16 Jul 2026 10:55:24 -0700 Subject: [PATCH 13/18] Clean up Signed-off-by: Hollow Man --- docs/design-docs/sparse-delta-refit.md | 5 +- .../generation/vllm/vllm_sparse_delta.py | 8 +- .../generation/vllm/vllm_sparse_refit.py | 61 ++---- nemo_rl/models/generation/vllm/vllm_worker.py | 4 +- .../workers/megatron_remote_sparse_refit.py | 8 +- nemo_rl/utils/weight_transfer_http.py | 14 +- nemo_rl/utils/weight_transfer_sparse_codec.py | 44 ++--- nemo_rl/utils/weight_transfer_stream.py | 144 ++++++-------- nemo_rl/utils/weight_transfer_zmq.py | 182 ++++++------------ .../vllm_remote_sparse_weight_synchronizer.py | 2 +- .../generation/test_vllm_sparse_refit.py | 49 ++--- .../unit/utils/test_weight_transfer_stream.py | 119 +++++------- ..._vllm_remote_sparse_weight_synchronizer.py | 1 - 13 files changed, 229 insertions(+), 412 deletions(-) diff --git a/docs/design-docs/sparse-delta-refit.md b/docs/design-docs/sparse-delta-refit.md index 707435fe227..15642f19b9d 100644 --- a/docs/design-docs/sparse-delta-refit.md +++ b/docs/design-docs/sparse-delta-refit.md @@ -172,8 +172,8 @@ S3 uses 64 MiB multipart parts, a 2 GiB CRT client memory limit, and a 10 Gbps throughput target. ZeroMQ assigns each producer to one inference-cluster relay. That root applies the payload locally and forwards it through a balanced binary relay tree, so each payload crosses the inter-cluster boundary once rather than -once per generation replica. Both transports use the same receiver endpoints -and checksum validation. +once per generation replica. The relay calls the same receiver decode/apply +queue as S3 directly, without a loopback HTTP data hop. Each ZeroMQ relay validates and deduplicates `(transfer_id, producer_id, payload_id)`, submits its local apply and at most two child forwards to bounded @@ -315,7 +315,6 @@ same nonempty token on producers and receivers. | `NRL_REFIT_ZMQ_SEND_WORKERS` | 4 | | `NRL_REFIT_ZMQ_RELAY_PAYLOAD_WORKERS` | 16 | | `NRL_REFIT_ZMQ_RELAY_FORWARD_WORKERS` | 8 | -| `NRL_REFIT_ZMQ_RELAY_FANOUT_WORKERS` | receiver endpoints x relay payload workers | | `NRL_REFIT_APPLY_QUEUE_DEPTH` / `NRL_REFIT_APPLY_BATCH_SIZE` | 32 / 8 | | `NRL_REFIT_PARTITION_WORKERS` | 2-8 from CPU count | | `NRL_REFIT_{S3,ZMQ}_ZSTD_THREADS` | 0 | diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py index d607e8a190e..1d77e7f388c 100644 --- a/nemo_rl/models/generation/vllm/vllm_sparse_delta.py +++ b/nemo_rl/models/generation/vllm/vllm_sparse_delta.py @@ -56,6 +56,7 @@ def __init__( self._targets = targets self._verification = verification self._source_storage = 0 + self._source_name = "" self._operation: sparse_codec.SparseOperation = "overwrite" self._sample_limit = 0 self._exact_sentinel: int | None = None @@ -69,11 +70,13 @@ def __init__( def start( self, + name: str, source: torch.Tensor, operation: sparse_codec.SparseOperation, sample_limit: int, exact_sentinel: int | None, ) -> None: + self._source_name = name self._source_storage = _storage_key(source) self._operation = operation self._sample_limit = sample_limit @@ -171,7 +174,8 @@ def __torch_dispatch__( if not xor_compatible: raise RuntimeError( - "XOR cannot pass through this native loader without changing semantics." + f"XOR for {self._source_name!r} cannot pass through this native " + "loader without changing semantics." ) source = source.expand_as(destination) destination_bits = sparse_codec.integer_view(destination) @@ -354,7 +358,7 @@ def observed_weights() -> Iterator[tuple[str, torch.Tensor]]: observations.append( (yielded_names[-1], mode.copies, mode.xor_compatible) ) - mode.start(source, operation, sample_limit, exact_sentinel) + mode.start(name, source, operation, sample_limit, exact_sentinel) active = True yielded_names.append(name) yield name, source diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py index f2c23c31a27..9965007eda9 100644 --- a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py +++ b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py @@ -19,6 +19,7 @@ import tempfile import threading import time +from collections.abc import Mapping from concurrent.futures import Future, ThreadPoolExecutor from typing import Any, Literal, NamedTuple, cast @@ -48,12 +49,6 @@ refit_env_int, ) from nemo_rl.utils.weight_transfer_zmq import ( - G_VLLM_REFIT_CHECKSUM_HEADER, - G_VLLM_REFIT_PAYLOAD_HEADER, - G_VLLM_REFIT_PRODUCER_HEADER, - G_VLLM_REFIT_TRANSFER_HEADER, - G_VLLM_REFIT_VERIFICATION_HEADER, - G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, ZmqSparseRefitServer, ) @@ -394,37 +389,22 @@ async def _apply_s3_manifest_payload( result["receiver_s3_download_s"] = download_s return result - async def _apply_zmq_payload(self, raw_request: Any) -> dict[str, Any]: - headers = raw_request.headers - transfer_id = headers.get(G_VLLM_REFIT_TRANSFER_HEADER, "") - producer_id = int(headers.get(G_VLLM_REFIT_PRODUCER_HEADER, "-1")) - payload_id = int(headers.get(G_VLLM_REFIT_PAYLOAD_HEADER, "-1")) - checksum = headers.get(G_VLLM_REFIT_CHECKSUM_HEADER, "") - verification_candidates = int( - headers.get(G_VLLM_REFIT_VERIFICATION_HEADER, "-1") - ) - if ( - not transfer_id - or producer_id < 0 - or payload_id < 0 - or not checksum - or verification_candidates < 0 - ): - raise ValueError("Missing or invalid ZeroMQ sparse refit payload headers.") - compressed = await raw_request.body() + def _apply_zmq_payload( + self, compressed: bytes, metadata: Mapping[str, Any] + ) -> dict[str, Any]: started = time.perf_counter() - payload = await asyncio.to_thread( - decode_sparse_payload, - compressed, - checksum, - ) + checksum = str(metadata["checksum"]) + payload = decode_sparse_payload(compressed, checksum) decode_s = time.perf_counter() - started - result = await asyncio.to_thread( - self._enqueue_sparse_payload_apply, + result = self._enqueue_sparse_payload_apply( payload, - (transfer_id, producer_id, payload_id), + ( + str(metadata["transfer_id"]), + int(metadata["producer_id"]), + int(metadata["payload_id"]), + ), checksum, - verification_candidates, + int(metadata["verification_candidates"]), ) result["receiver_zmq_decode_s"] = decode_s return result @@ -435,7 +415,7 @@ def setup_api_server(self, app: Any) -> None: async def respond( raw_request: Request, - action: Literal["prepare", "s3", "flush", "zmq", "zmq_flush"], + action: Literal["prepare", "s3", "flush", "zmq_flush"], ) -> JSONResponse: if cfg["vllm_cfg"]["async_engine"]: self._refit_async_loop = asyncio.get_running_loop() @@ -456,8 +436,6 @@ async def respond( result = await self._apply_s3_manifest_payload( await raw_request.json() ) - elif action == "zmq": - result = await self._apply_zmq_payload(raw_request) elif action == "zmq_flush": body = await raw_request.json() result = await asyncio.to_thread( @@ -475,7 +453,7 @@ async def respond( ) def endpoint( - action: Literal["prepare", "s3", "flush", "zmq", "zmq_flush"], + action: Literal["prepare", "s3", "flush", "zmq_flush"], ): async def handle(raw_request: Request) -> JSONResponse: return await respond(raw_request, action) @@ -486,7 +464,6 @@ async def handle(raw_request: Request) -> JSONResponse: (G_VLLM_REFIT_S3_MANIFEST_PATH, "s3"), (G_VLLM_REFIT_PREPARE_PATH, "prepare"), (G_VLLM_REFIT_FLUSH_PATH, "flush"), - (G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, "zmq"), (G_VLLM_REFIT_ZMQ_FLUSH_PATH, "zmq_flush"), ): app.add_api_route(path, endpoint(action), methods=["POST"]) @@ -497,7 +474,7 @@ def report_refit_server_base_url(self) -> str | None: base_url = getattr(self._worker, "base_url", None) return base_url.removesuffix("/v1") if base_url else None - def start_zmq_sparse_refit_relay(self, refit_urls: list[str]) -> str: + def start_zmq_sparse_refit_relay(self) -> str: if self._zmq_refit_server is not None: return self._zmq_refit_server[1] cfg = self._worker.cfg @@ -506,7 +483,7 @@ def start_zmq_sparse_refit_relay(self, refit_urls: list[str]) -> str: cfg.get("port_range_high", DEFAULT_GENERATION_PORT_RANGE_HIGH), ) server = ZmqSparseRefitServer( - refit_urls, + self._apply_zmq_payload, bind_address=f"tcp://0.0.0.0:{port}", api_key_env_var=cfg["vllm_cfg"].get("http_refit_api_key_env_var"), timeout_s=float(os.getenv("NRL_REFIT_ZMQ_TIMEOUT_S") or 600.0), @@ -520,14 +497,10 @@ def start_zmq_sparse_refit_relay(self, refit_urls: list[str]) -> str: def configure_zmq_sparse_refit_relay(self, relay_addresses: list[str]) -> None: if self._zmq_refit_server is None: raise RuntimeError("ZeroMQ sparse refit relay is not running.") - local_refit_url = self.report_refit_server_base_url() - if local_refit_url is None: - raise RuntimeError("Local vLLM sparse refit endpoint is unavailable.") server, own_address = self._zmq_refit_server server.configure_tree( relay_addresses, own_address=own_address, - local_refit_url=local_refit_url, ) def stop_zmq_sparse_refit_relay(self) -> None: diff --git a/nemo_rl/models/generation/vllm/vllm_worker.py b/nemo_rl/models/generation/vllm/vllm_worker.py index 316f8c6bbf3..b3a8aa2b63d 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker.py +++ b/nemo_rl/models/generation/vllm/vllm_worker.py @@ -662,11 +662,11 @@ def report_refit_server_base_url(self) -> str | None: receiver = self._sparse_refit_receiver return receiver.report_refit_server_base_url() if receiver is not None else None - def start_zmq_sparse_refit_relay(self, refit_urls: list[str]) -> str: + def start_zmq_sparse_refit_relay(self) -> str: receiver = self._sparse_refit_receiver if receiver is None: raise RuntimeError("Remote sparse refit is not enabled for this worker.") - return receiver.start_zmq_sparse_refit_relay(refit_urls) + return receiver.start_zmq_sparse_refit_relay() def configure_zmq_sparse_refit_relay(self, relay_addresses: list[str]) -> None: receiver = self._sparse_refit_receiver diff --git a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py index 3da8550596f..2b86eba287d 100644 --- a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py +++ b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py @@ -14,7 +14,6 @@ """Canonical Hugging Face sparse-refit state for a Megatron policy worker.""" -from collections.abc import Iterator from typing import Any import torch @@ -33,9 +32,6 @@ def __init__(self, worker: Any, delta_config: VllmDeltaCompressionConfig) -> Non self._worker = worker self._tracker = DeltaCompressionTracker(delta_config.model_dump()) - def _iter_params(self) -> Iterator[tuple[str, torch.Tensor]]: - return self._worker._iter_params_with_optional_kv_scales() - def initialize_baseline( self, *, @@ -44,7 +40,7 @@ def initialize_baseline( transport: str, ) -> dict[str, tuple[tuple[int, ...], torch.dtype]]: init_sparse_delta_baseline_from_iterator( - self._iter_params(), + self._worker._iter_params_with_optional_kv_scales(), delta_tracker=self._tracker, shard_rank=shard_rank, shard_count=shard_count, @@ -73,7 +69,7 @@ def stream( }[transport] self._tracker.overwrite_names = frozenset(overwrite_names) result = streamer( - self._iter_params(), + self._worker._iter_params_with_optional_kv_scales(), delta_tracker=self._tracker, transfer_id=transfer_id, refit_targets=targets, diff --git a/nemo_rl/utils/weight_transfer_http.py b/nemo_rl/utils/weight_transfer_http.py index 29168b2e8da..685bc9349d9 100644 --- a/nemo_rl/utils/weight_transfer_http.py +++ b/nemo_rl/utils/weight_transfer_http.py @@ -79,24 +79,19 @@ def _http_executor(workers: int) -> ThreadPoolExecutor: def post_vllm_refit_endpoints( endpoint_urls: Sequence[str], - body: Mapping[str, Any] | bytes, + body: Mapping[str, Any], *, api_key: str | None, timeout_s: float, - headers: Mapping[str, str] | None = None, - executor: ThreadPoolExecutor | None = None, ) -> list[dict[str, Any]]: - request_headers = dict(headers or {}) + request_headers = {} if api_key: request_headers[G_VLLM_REFIT_API_KEY_HEADER] = api_key - request_kwargs: dict[str, Any] = ( - {"data": body} if isinstance(body, bytes) else {"json": body} - ) def post(url: str) -> dict[str, Any]: response = refit_http_session().post( url, - **request_kwargs, + json=body, headers=request_headers, timeout=timeout_s, ) @@ -111,8 +106,7 @@ def post(url: str) -> dict[str, Any]: ) return result - pool = executor or _http_executor(len(endpoint_urls)) - return list(pool.map(post, endpoint_urls)) + return list(_http_executor(len(endpoint_urls)).map(post, endpoint_urls)) def merge_vllm_refit_metrics( diff --git a/nemo_rl/utils/weight_transfer_sparse_codec.py b/nemo_rl/utils/weight_transfer_sparse_codec.py index c53216c441f..4a9d5ecda3b 100644 --- a/nemo_rl/utils/weight_transfer_sparse_codec.py +++ b/nemo_rl/utils/weight_transfer_sparse_codec.py @@ -104,14 +104,6 @@ def integer_view(tensor: torch.Tensor) -> torch.Tensor: return tensor.view(integer_dtype_for_element_size(tensor.element_size())) -def _bytewise_diff_mask(current: torch.Tensor, baseline: torch.Tensor) -> torch.Tensor: - if current.shape != baseline.shape or current.dtype != baseline.dtype: - raise ValueError( - "Current tensor and baseline must have identical shape and dtype." - ) - return integer_view(current) != integer_view(baseline) - - def encode_sparse_infos( infos: Iterable[SparseInfo], ) -> TensorPayload: @@ -241,15 +233,23 @@ def __init__( def prepare_sparse_delta_payload( self, tensors: TensorBatch ) -> PreparedTensorPayload: + self._wait_for_baseline_commits() sparse_infos: list[SparseInfo] = [] + pending_updates = {} changed_elements = total_elements = 0 - for name, baseline, current, locations, current_values in self._changes( - tensors - ): + for name, tensor in tensors: + baseline = self.baseline.get(name) + if baseline is None: + raise RuntimeError(f"Sparse delta baseline is missing {name!r}.") + current = tensor.detach().cpu().contiguous() baseline_bits = integer_view(baseline).view(-1) + current_bits = integer_view(current).view(-1) + locations = current_bits.ne(baseline_bits).nonzero().view(-1) total_elements += current.numel() changed_elements += locations.numel() if locations.numel(): + current_values = current_bits[locations] + pending_updates[name] = (locations, current_values) operation: SparseOperation = ( "overwrite" if name in self.overwrite_names else self.encoding ) @@ -267,31 +267,13 @@ def prepare_sparse_delta_payload( operation, ) ) + with self._pending_updates_lock: + self._pending_updates.update(pending_updates) payload = encode_sparse_infos(sparse_infos) if self.verification_samples: self._add_verification_samples(payload[2]) return payload, changed_elements, total_elements - def _changes( - self, tensors: Iterable[NamedTensor] - ) -> Iterable[tuple[str, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]]: - self._wait_for_baseline_commits() - pending_updates = {} - for name, tensor in tensors: - baseline = self.baseline.get(name) - if baseline is None: - raise RuntimeError(f"Sparse delta baseline is missing {name!r}.") - current = tensor.detach().cpu().contiguous() - locations = ( - _bytewise_diff_mask(current, baseline).view(-1).nonzero().view(-1) - ) - values = integer_view(current).view(-1)[locations] - if locations.numel(): - pending_updates[name] = (locations, values) - yield name, baseline, current, locations, values - with self._pending_updates_lock: - self._pending_updates.update(pending_updates) - def _add_verification_samples( self, metadata: list[dict[str, Any]], diff --git a/nemo_rl/utils/weight_transfer_stream.py b/nemo_rl/utils/weight_transfer_stream.py index 3bcd00f1839..2645f884454 100644 --- a/nemo_rl/utils/weight_transfer_stream.py +++ b/nemo_rl/utils/weight_transfer_stream.py @@ -24,7 +24,7 @@ from contextlib import suppress from dataclasses import dataclass from functools import cache -from typing import Any, Protocol +from typing import Any, Callable from urllib.parse import quote import torch @@ -51,19 +51,14 @@ _S3_MEMORY_LIMIT = 2 * 1024**3 -class SparseRefitTransport(Protocol): - """Payload delivery owned by one generic stream invocation.""" +@dataclass(frozen=True) +class SparseRefitTransport: + """Transport-specific callbacks used by the shared streaming pipeline.""" name: str transfer_workers: int - - def send( - self, body: bytes, payload_id: int, verification_candidates: int - ) -> dict[str, Any]: ... - - def cleanup(self) -> None: - """Release resources created on the current transfer worker.""" - ... + send: Callable[[bytes, int, int], dict[str, Any]] + cleanup: Callable[[], None] @dataclass @@ -170,75 +165,6 @@ def _request(self, method: str, key: str, body: bytes | None = None) -> Any: ) -class _S3ManifestTransport: - name = "s3" - - def __init__( - self, - *, - store: _S3ObjectStore, - endpoint_urls: Sequence[str], - run_prefix: str, - api_key: str | None, - timeout_s: float, - transfer_workers: int, - ) -> None: - self._store = store - self._endpoint_urls = endpoint_urls - self._run_prefix = run_prefix - self._api_key = api_key - self._timeout_s = timeout_s - self.transfer_workers = transfer_workers - self._keys = threading.local() - - def send( - self, body: bytes, payload_id: int, verification_candidates: int - ) -> dict[str, Any]: - key = f"{self._run_prefix}/{payload_id:06d}.pt" - keys = getattr(self._keys, "values", None) - if keys is None: - keys = [] - self._keys.values = keys - keys.append(key) - - try: - started = time.perf_counter() - self._store.put(key, body) - s3_put_s = time.perf_counter() - started - started = time.perf_counter() - responses = post_vllm_refit_endpoints( - self._endpoint_urls, - { - "bucket": self._store.bucket, - "region": self._store.region, - "key": key, - "checksum": sparse_payload_checksum(body), - "verification_candidates": verification_candidates, - }, - api_key=self._api_key, - timeout_s=self._timeout_s, - ) - result = { - "s3_put_s": s3_put_s, - "manifest_post_s": time.perf_counter() - started, - "receiver": merge_vllm_refit_metrics({}, responses, maximum=True), - } - finally: - try: - self._store.delete(key) - except Exception: - pass - else: - keys.remove(key) - return result - - def cleanup(self) -> None: - for key in getattr(self._keys, "values", ()): - with suppress(Exception): - self._store.delete(key) - self._keys.values = [] - - def refit_env_int(name: str, *, default: int, min_value: int = 1) -> int: value = int(os.getenv(name) or default) if value < min_value: @@ -611,19 +537,65 @@ def stream_sparse_delta_payloads_via_s3_manifest( if object_prefix else f"{transfer_id}/{shard_rank:06d}" ) + api_key = vllm_refit_api_key(api_key_env_var) + keys = threading.local() + + def send( + body: bytes, payload_id: int, verification_candidates: int + ) -> dict[str, Any]: + key = f"{run_prefix}/{payload_id:06d}.pt" + pending = getattr(keys, "values", None) + if pending is None: + pending = keys.values = [] + pending.append(key) + try: + started = time.perf_counter() + store.put(key, body) + s3_put_s = time.perf_counter() - started + started = time.perf_counter() + responses = post_vllm_refit_endpoints( + endpoint_urls, + { + "bucket": store.bucket, + "region": store.region, + "key": key, + "checksum": sparse_payload_checksum(body), + "verification_candidates": verification_candidates, + }, + api_key=api_key, + timeout_s=timeout_s, + ) + result = { + "s3_put_s": s3_put_s, + "manifest_post_s": time.perf_counter() - started, + "receiver": merge_vllm_refit_metrics({}, responses, maximum=True), + } + finally: + try: + store.delete(key) + except Exception: + pass + else: + pending.remove(key) + return result + + def cleanup() -> None: + for key in getattr(keys, "values", ()): + with suppress(Exception): + store.delete(key) + keys.values = [] + return stream_sparse_delta_payloads( iterator, delta_tracker=delta_tracker, - transport=_S3ManifestTransport( - store=store, - endpoint_urls=endpoint_urls, - run_prefix=run_prefix, - api_key=vllm_refit_api_key(api_key_env_var), - timeout_s=timeout_s, + transport=SparseRefitTransport( + name="s3", transfer_workers=refit_env_int( "NRL_REFIT_S3_UPLOAD_WORKERS", default=max(4, min(32, os.cpu_count() or 32)), ), + send=send, + cleanup=cleanup, ), shard_rank=shard_rank, shard_count=shard_count, diff --git a/nemo_rl/utils/weight_transfer_zmq.py b/nemo_rl/utils/weight_transfer_zmq.py index fb00b998955..1ec44ebacef 100644 --- a/nemo_rl/utils/weight_transfer_zmq.py +++ b/nemo_rl/utils/weight_transfer_zmq.py @@ -18,7 +18,7 @@ import threading import time import uuid -from collections.abc import Iterable, Mapping, Sequence +from collections.abc import Callable, Iterable, Mapping, Sequence from concurrent.futures import Future, ThreadPoolExecutor from contextlib import suppress from dataclasses import dataclass, field @@ -28,27 +28,19 @@ from nemo_rl.utils.weight_transfer_http import ( merge_vllm_refit_metrics, - post_vllm_refit_endpoints, vllm_refit_api_key, - vllm_refit_endpoints, ) from nemo_rl.utils.weight_transfer_sparse_codec import ( DeltaCompressionTracker, NamedTensor, ) from nemo_rl.utils.weight_transfer_stream import ( + SparseRefitTransport, refit_env_int, sparse_payload_checksum, stream_sparse_delta_payloads, ) -G_VLLM_REFIT_ZMQ_PAYLOAD_PATH = "/nemo-rl/refit/zmq-payload" -G_VLLM_REFIT_TRANSFER_HEADER = "x-nemo-rl-refit-transfer" -G_VLLM_REFIT_PRODUCER_HEADER = "x-nemo-rl-refit-producer" -G_VLLM_REFIT_PAYLOAD_HEADER = "x-nemo-rl-refit-payload" -G_VLLM_REFIT_CHECKSUM_HEADER = "x-nemo-rl-refit-checksum" -G_VLLM_REFIT_VERIFICATION_HEADER = "x-nemo-rl-refit-verification-candidates" - _PROTOCOL = "nemo-rl-sparse-zmq-v1" _DATA = b"DATA" _ACK = b"ACK" @@ -170,75 +162,18 @@ def close(self) -> None: self._socket.close() -class _ZmqTransport: - name = "zmq" - - def __init__( - self, - address: str, - *, - transfer_id: str, - timeout_s: float, - producer_id: int, - api_key: str | None, - transfer_workers: int, - ) -> None: - self._address = address - self._transfer_id = transfer_id - self._timeout_s = timeout_s - self._producer_id = producer_id - self._api_key = api_key - self.transfer_workers = transfer_workers - self._local = threading.local() - - def send( - self, body: bytes, payload_id: int, verification_candidates: int - ) -> dict[str, Any]: - client = getattr(self._local, "client", None) - if client is None: - client = ZmqSparseRefitClient( - self._address, - timeout_s=self._timeout_s, - producer_id=self._producer_id, - api_key=self._api_key, - ) - self._local.client = client - started = time.perf_counter() - reply = client.send_payload( - transfer_id=self._transfer_id, - payload_id=payload_id, - checksum=sparse_payload_checksum(body), - verification_candidates=verification_candidates, - body=body, - ) - return { - "zmq_send_s": time.perf_counter() - started, - "receiver": reply, - } - - def cleanup(self) -> None: - client = getattr(self._local, "client", None) - if client is not None: - client.close() - del self._local.client - - class ZmqSparseRefitServer: - """Bounded ROUTER relay that fans each compressed payload to all replicas.""" + """Bounded ROUTER relay that applies locally and fans out through a tree.""" def __init__( self, - refit_urls: Sequence[str], + apply_payload: Callable[[bytes, Mapping[str, Any]], dict[str, Any]], *, bind_address: str, api_key_env_var: str | None, timeout_s: float, ) -> None: - self._refit_endpoints = vllm_refit_endpoints( - refit_urls, G_VLLM_REFIT_ZMQ_PAYLOAD_PATH - ) - if not self._refit_endpoints: - raise ValueError("ZeroMQ sparse refit requires receiver HTTP URLs.") + self._apply_payload = apply_payload self._bind_address = bind_address self._token = vllm_refit_api_key(api_key_env_var) self._timeout_s = timeout_s @@ -250,15 +185,11 @@ def __init__( self._payload_workers = refit_env_int( "NRL_REFIT_ZMQ_RELAY_PAYLOAD_WORKERS", default=16 ) - self._fanout_workers = refit_env_int( - "NRL_REFIT_ZMQ_RELAY_FANOUT_WORKERS", - default=len(self._refit_endpoints) * self._payload_workers, - ) self._transfer_lock = threading.Lock() self._transfer_condition = threading.Condition(self._transfer_lock) self._transfers: dict[str, _RelayTransfer] = {} self._flush_results: dict[str, Future[dict[str, Any]]] = {} - self._tree: tuple[tuple[str, ...], int, str] | None = None + self._tree: tuple[tuple[str, ...], int] | None = None self._forward_local = threading.local() self._forward_clients: list[ZmqSparseRefitClient] = [] self._payload_executor = ThreadPoolExecutor( @@ -269,23 +200,15 @@ def __init__( max_workers=refit_env_int("NRL_REFIT_ZMQ_RELAY_FORWARD_WORKERS", default=8), thread_name_prefix="nrl-zmq-forward", ) - self._http_executor = ThreadPoolExecutor( - max_workers=self._fanout_workers, - thread_name_prefix="nrl-zmq-fanout", - ) def configure_tree( self, relay_addresses: Sequence[str], *, own_address: str, - local_refit_url: str, ) -> None: addresses = tuple(dict.fromkeys(relay_addresses)) - (local_endpoint,) = vllm_refit_endpoints( - [local_refit_url], G_VLLM_REFIT_ZMQ_PAYLOAD_PATH - ) - self._tree = (addresses, addresses.index(own_address), local_endpoint) + self._tree = (addresses, addresses.index(own_address)) def start(self) -> str: self._thread = threading.Thread( @@ -317,10 +240,12 @@ def flush(self, transfer_id: str, expected_payloads: int = 0) -> dict[str, Any]: completion = self._flush_results.get(transfer_id) if completion is None: ready = self._transfer_condition.wait_for( - lambda: len( - self._transfers.get(transfer_id, _RelayTransfer()).checksums - ) - >= expected_payloads, + lambda: ( + len( + self._transfers.get(transfer_id, _RelayTransfer()).checksums + ) + >= expected_payloads + ), timeout=self._timeout_s, ) if not ready: @@ -371,29 +296,10 @@ def _fanout( body: bytes, metadata: Mapping[str, Any], ) -> dict[str, Any]: - headers = { - "content-type": "application/octet-stream", - G_VLLM_REFIT_TRANSFER_HEADER: str(metadata["transfer_id"]), - G_VLLM_REFIT_PRODUCER_HEADER: str(metadata["producer_id"]), - G_VLLM_REFIT_PAYLOAD_HEADER: str(metadata["payload_id"]), - G_VLLM_REFIT_CHECKSUM_HEADER: str(metadata["checksum"]), - G_VLLM_REFIT_VERIFICATION_HEADER: str(metadata["verification_candidates"]), - } started = time.perf_counter() - endpoints = self._refit_endpoints - if self._tree is not None: - endpoints = [self._tree[2]] - results = post_vllm_refit_endpoints( - endpoints, - body, - api_key=self._token, - timeout_s=self._timeout_s, - headers=headers, - executor=self._http_executor, - ) - merged = merge_vllm_refit_metrics({}, results, maximum=True) - merged["receiver_relay_fanout_s"] = time.perf_counter() - started - return merged + result = self._apply_payload(body, metadata) + result["receiver_relay_fanout_s"] = time.perf_counter() - started + return result def _forward( self, @@ -459,7 +365,14 @@ def _parse_data_message( producer_id = int(metadata["producer_id"]) payload_id = int(metadata["payload_id"]) checksum = str(metadata["checksum"]) - if not transfer_id or producer_id < 0 or payload_id < 0: + verification_candidates = int(metadata["verification_candidates"]) + if ( + not transfer_id + or producer_id < 0 + or payload_id < 0 + or not checksum + or verification_candidates < 0 + ): raise ValueError("Invalid ZeroMQ sparse refit payload identity.") relay_root = metadata.get("relay_root") if relay_root is not None and ( @@ -508,7 +421,7 @@ def _run(self) -> None: ) ) if self._tree is not None: - addresses, own_index, _ = self._tree + addresses, own_index = self._tree relay_root = str( metadata.get("relay_root", addresses[own_index]) ) @@ -560,7 +473,6 @@ def _run(self) -> None: finally: self._payload_executor.shutdown(wait=True, cancel_futures=True) self._forward_executor.shutdown(wait=True, cancel_futures=True) - self._http_executor.shutdown(wait=True, cancel_futures=True) for client in self._forward_clients: client.close() socket.close() @@ -582,16 +494,48 @@ def stream_sparse_delta_payloads_via_zmq( if not addresses: raise ValueError("At least one ZeroMQ sparse refit address is required.") address = addresses[shard_rank % len(addresses)] + api_key = vllm_refit_api_key(api_key_env_var) + local = threading.local() + + def send( + body: bytes, payload_id: int, verification_candidates: int + ) -> dict[str, Any]: + client = getattr(local, "client", None) + if client is None: + client = ZmqSparseRefitClient( + address, + timeout_s=timeout_s, + producer_id=shard_rank, + api_key=api_key, + ) + local.client = client + started = time.perf_counter() + reply = client.send_payload( + transfer_id=transfer_id, + payload_id=payload_id, + checksum=sparse_payload_checksum(body), + verification_candidates=verification_candidates, + body=body, + ) + return { + "zmq_send_s": time.perf_counter() - started, + "receiver": reply, + } + + def cleanup() -> None: + client = getattr(local, "client", None) + if client is not None: + client.close() + del local.client + return stream_sparse_delta_payloads( iterator, delta_tracker=delta_tracker, - transport=_ZmqTransport( - address, - transfer_id=transfer_id, - timeout_s=timeout_s, - producer_id=shard_rank, - api_key=vllm_refit_api_key(api_key_env_var), + transport=SparseRefitTransport( + name="zmq", transfer_workers=refit_env_int("NRL_REFIT_ZMQ_SEND_WORKERS", default=4), + send=send, + cleanup=cleanup, ), shard_rank=shard_rank, shard_count=shard_count, diff --git a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py index b55222bb600..66936b06c38 100644 --- a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py +++ b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py @@ -341,7 +341,7 @@ def init_communicator(self) -> None: self._targets = [ address for address in self._run_generation_workers( - "start_zmq_sparse_refit_relay", refit_urls=self._refit_urls + "start_zmq_sparse_refit_relay" ) if address ] diff --git a/tests/unit/models/generation/test_vllm_sparse_refit.py b/tests/unit/models/generation/test_vllm_sparse_refit.py index bb766a221eb..e0e82d1b645 100644 --- a/tests/unit/models/generation/test_vllm_sparse_refit.py +++ b/tests/unit/models/generation/test_vllm_sparse_refit.py @@ -46,14 +46,6 @@ encode_sparse_infos, sparse_locations_for_item, ) -from nemo_rl.utils.weight_transfer_zmq import ( - G_VLLM_REFIT_CHECKSUM_HEADER, - G_VLLM_REFIT_PAYLOAD_HEADER, - G_VLLM_REFIT_PRODUCER_HEADER, - G_VLLM_REFIT_TRANSFER_HEADER, - G_VLLM_REFIT_VERIFICATION_HEADER, - G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, -) @contextmanager @@ -422,21 +414,16 @@ async def test_sparse_refit_payload_handlers_decode_and_enqueue(monkeypatch) -> ) assert s3_result["ok"] - invalid_request = SimpleNamespace(headers={}, body=AsyncMock()) - with pytest.raises(ValueError, match="Missing or invalid"): - await receiver._apply_zmq_payload(invalid_request) - - request = SimpleNamespace( - headers={ - G_VLLM_REFIT_TRANSFER_HEADER: "transfer", - G_VLLM_REFIT_PRODUCER_HEADER: "2", - G_VLLM_REFIT_PAYLOAD_HEADER: "3", - G_VLLM_REFIT_CHECKSUM_HEADER: "checksum", - G_VLLM_REFIT_VERIFICATION_HEADER: "5", + zmq_result = receiver._apply_zmq_payload( + b"compressed", + { + "transfer_id": "transfer", + "producer_id": 2, + "payload_id": 3, + "checksum": "checksum", + "verification_candidates": 5, }, - body=AsyncMock(return_value=b"compressed"), ) - zmq_result = await receiver._apply_zmq_payload(request) assert zmq_result["ok"] assert enqueue.call_args_list == [ call(b"s3-payload", ("object-key", -1, -1), "checksum", 4), @@ -459,9 +446,6 @@ def test_sparse_refit_api_auth_dispatch_and_error_mapping(monkeypatch) -> None: receiver._apply_s3_manifest_payload = AsyncMock( return_value={"ok": True, "payloads": 1} ) - receiver._apply_zmq_payload = AsyncMock( - side_effect=RuntimeError("apply failed") - ) receiver._flush_queued_sparse_payloads = MagicMock( return_value={"ok": True, "payloads": 2} ) @@ -491,21 +475,13 @@ def test_sparse_refit_api_auth_dispatch_and_error_mapping(monkeypatch) -> None: json={"tensors": {"weight": [[2, 3], "bfloat16"]}}, headers=headers, ) - zmq_response = client.post( - G_VLLM_REFIT_ZMQ_PAYLOAD_PATH, - content=b"payload", - headers=headers, - ) - assert unauthorized.status_code == 403 assert s3_response.status_code == 200 assert flush_response.status_code == 200 assert zmq_flush_response.status_code == 200 assert prepare_response.status_code == 200 - assert zmq_response.status_code == 500 assert receiver._refit_async_loop is not None receiver._apply_s3_manifest_payload.assert_awaited_once_with({"key": "key"}) - receiver._apply_zmq_payload.assert_awaited_once() receiver._flush_queued_sparse_payloads.assert_called_once_with() receiver.flush_zmq_sparse_refit_relay.assert_called_once_with("transfer", 0) receiver._refit_collective_rpc.assert_called_once_with( @@ -611,8 +587,12 @@ def test_async_sparse_refit_exposes_zmq_relay(monkeypatch) -> None: with _sparse_refit_receiver(async_engine=True, config=config) as receiver: receiver._worker.base_url = "http://10.0.0.1:8000/v1" - assert receiver.start_zmq_sparse_refit_relay(["http://receiver"]) == ( - "tcp://10.0.0.1:12345" + assert receiver.start_zmq_sparse_refit_relay() == "tcp://10.0.0.1:12345" + server_type.assert_called_once_with( + receiver._apply_zmq_payload, + bind_address="tcp://0.0.0.0:12345", + api_key_env_var=None, + timeout_s=600.0, ) server.start.assert_called_once_with() receiver.configure_zmq_sparse_refit_relay( @@ -621,7 +601,6 @@ def test_async_sparse_refit_exposes_zmq_relay(monkeypatch) -> None: server.configure_tree.assert_called_once_with( ["tcp://10.0.0.1:12345", "tcp://10.0.0.2:12345"], own_address="tcp://10.0.0.1:12345", - local_refit_url="http://10.0.0.1:8000", ) server.flush.return_value = {"ok": True, "payloads": 2} assert receiver.flush_zmq_sparse_refit_relay("transfer") == { diff --git a/tests/unit/utils/test_weight_transfer_stream.py b/tests/unit/utils/test_weight_transfer_stream.py index 35c68de85b1..9acc9885b51 100644 --- a/tests/unit/utils/test_weight_transfer_stream.py +++ b/tests/unit/utils/test_weight_transfer_stream.py @@ -16,7 +16,6 @@ import json import threading from concurrent.futures import ThreadPoolExecutor -from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from types import SimpleNamespace import pytest @@ -35,16 +34,12 @@ sparse_locations_for_item, ) from nemo_rl.utils.weight_transfer_stream import ( + SparseRefitTransport, download_s3_refit_payload, sparse_export_chunk_size, sparse_payload_checksum, ) from nemo_rl.utils.weight_transfer_zmq import ( - G_VLLM_REFIT_CHECKSUM_HEADER, - G_VLLM_REFIT_PAYLOAD_HEADER, - G_VLLM_REFIT_PRODUCER_HEADER, - G_VLLM_REFIT_TRANSFER_HEADER, - G_VLLM_REFIT_VERIFICATION_HEADER, ZmqSparseRefitClient, ZmqSparseRefitServer, ) @@ -81,23 +76,14 @@ def snapshot_baseline(self, chunk) -> None: self.names.extend(name for name, _tensor in chunk) -class _SparseTestTransport: - name = "zmq" - transfer_workers = 1 - - def __init__(self, send_payload) -> None: - self._send_payload = send_payload - self.cleaned = False - - def send(self, body, payload_id, _verification_candidates): - return self._send_payload(body, payload_id) - - def cleanup(self) -> None: - self.cleaned = True - - def _stream_sparse_test_payloads(tensors, send_payload): - transport = _SparseTestTransport(send_payload) + cleaned = [] + transport = SparseRefitTransport( + name="zmq", + transfer_workers=1, + send=lambda body, payload_id, _candidates: send_payload(body, payload_id), + cleanup=lambda: cleaned.append(True), + ) try: return weight_transfer_stream.stream_sparse_delta_payloads( tensors, @@ -107,7 +93,7 @@ def _stream_sparse_test_payloads(tensors, send_payload): shard_count=1, ) finally: - assert transport.cleaned + assert cleaned def _delta_tracker(encoding: str = "overwrite") -> DeltaCompressionTracker: @@ -434,7 +420,13 @@ def send(body, payload_index): ) return {"receiver": {}} - transport = _SparseTestTransport(send) + cleaned = [] + transport = SparseRefitTransport( + "zmq", + 1, + lambda body, payload_id, _candidates: send(body, payload_id), + lambda: cleaned.append(True), + ) result = weight_transfer_stream.stream_sparse_delta_payloads( tensors, delta_tracker=tracker, @@ -443,7 +435,7 @@ def send(body, payload_index): shard_count=1, ) - assert transport.cleaned + assert cleaned assert result == {"payloads": 2, "changed_elements": 4, "total_elements": 4} payload_names = [ [item["name"] for item in payloads[index][2]] for index in sorted(payloads) @@ -485,8 +477,13 @@ def test_sparse_baseline_snapshots_only_owned_export_chunks( def test_sparse_stream_sends_only_owned_export_chunks(monkeypatch) -> None: monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "1") sent = [] - transport = _SparseTestTransport( - lambda _body, payload_id: sent.append(payload_id) or {"receiver": {}} + transport = SparseRefitTransport( + "zmq", + 1, + lambda _body, payload_id, _candidates: ( + sent.append(payload_id) or {"receiver": {}} + ), + lambda: None, ) result = weight_transfer_stream.stream_sparse_delta_payloads( @@ -759,6 +756,8 @@ def frames(kind: bytes = b"DATA", **updates: object) -> list[bytes]: (frames(protocol="other"), ValueError, "protocol"), (frames(api_key="wrong"), PermissionError, "authentication"), (frames(transfer_id=""), ValueError, "identity"), + (frames(checksum=""), ValueError, "identity"), + (frames(verification_candidates=-1), ValueError, "identity"), ): with pytest.raises(error, match=match): server._parse_data_message(message) @@ -766,34 +765,6 @@ def frames(kind: bytes = b"DATA", **updates: object) -> list[bytes]: assert server._parse_data_message(frames())[1] == ("transfer", 0, 1) -def _receiver_server(received): - class Handler(BaseHTTPRequestHandler): - def do_POST(self): - body = self.rfile.read(int(self.headers["content-length"])) - received.append( - ({key.lower(): value for key, value in self.headers.items()}, body) - ) - ok = self.headers[G_VLLM_REFIT_CHECKSUM_HEADER] == sparse_payload_checksum( - body - ) - response = json.dumps( - {"ok": ok, "receiver_total_s": 0.25, "error": "checksum mismatch"} - ).encode() - self.send_response(200 if ok else 500) - self.send_header("content-type", "application/json") - self.send_header("content-length", str(len(response))) - self.end_headers() - self.wfile.write(response) - - def log_message(self, *args): - pass - - server = ThreadingHTTPServer(("127.0.0.1", 0), Handler) - thread = threading.Thread(target=server.serve_forever, daemon=True) - thread.start() - return server, thread - - def _send_zmq_payload( client: ZmqSparseRefitClient, payload_id: int, @@ -813,20 +784,28 @@ def _send_zmq_payload( def test_zmq_sparse_refit_relay_fans_out_and_rejects_corruption(monkeypatch) -> None: monkeypatch.setenv("NRL_TEST_REFIT_KEY", "secret") received = [[] for _ in range(4)] - receivers = [_receiver_server(items) for items in received] - urls = [f"http://127.0.0.1:{server.server_port}" for server, _ in receivers] + + def apply(items): + def apply_payload(body, metadata): + if metadata["checksum"] != sparse_payload_checksum(body): + raise ValueError("checksum mismatch") + items.append((dict(metadata), body)) + return {"ok": True, "receiver_total_s": 0.25} + + return apply_payload + relays = [ ZmqSparseRefitServer( - urls, + apply(items), bind_address="tcp://127.0.0.1:*", api_key_env_var="NRL_TEST_REFIT_KEY", timeout_s=5.0, ) - for _ in urls + for items in received ] addresses = [relay.start() for relay in relays] - for relay, address, url in zip(relays, addresses, urls, strict=True): - relay.configure_tree(addresses, own_address=address, local_refit_url=url) + for relay, address in zip(relays, addresses, strict=True): + relay.configure_tree(addresses, own_address=address) unauthenticated_client = ZmqSparseRefitClient( addresses[0], timeout_s=5.0, @@ -853,14 +832,14 @@ def test_zmq_sparse_refit_relay_fans_out_and_rejects_corruption(monkeypatch) -> assert flushed["receiver_total_s"] == 0.25 assert [len(items) for items in received] == [1] * 4 for items in received: - headers, posted_body = items[0] + metadata, posted_body = items[0] assert posted_body == body - assert headers[G_VLLM_REFIT_TRANSFER_HEADER] == "transfer-a" - assert headers[G_VLLM_REFIT_PRODUCER_HEADER] == "3" - assert headers[G_VLLM_REFIT_PAYLOAD_HEADER] == "7" - assert headers[G_VLLM_REFIT_CHECKSUM_HEADER] == checksum - assert headers[G_VLLM_REFIT_VERIFICATION_HEADER] == "2" - assert headers["x-nemo-rl-refit-key"] == "secret" + assert metadata["transfer_id"] == "transfer-a" + assert metadata["producer_id"] == 3 + assert metadata["payload_id"] == 7 + assert metadata["checksum"] == checksum + assert metadata["verification_candidates"] == 2 + assert metadata["api_key"] == "secret" with pytest.raises(RuntimeError, match="already flushed"): _send_zmq_payload(client, 8, body) @@ -872,7 +851,3 @@ def test_zmq_sparse_refit_relay_fans_out_and_rejects_corruption(monkeypatch) -> client.close() for relay in relays: relay.close() - for server, thread in receivers: - server.shutdown() - thread.join(timeout=5.0) - server.server_close() 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 c54c13696d3..c7078ac91fd 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 @@ -252,7 +252,6 @@ def test_initializes_streams_commits_and_updates_baseline( generation.worker_group.run_all_workers_single_data.assert_any_call( "start_zmq_sparse_refit_relay", run_rank_0_only_axes=["tensor_parallel", "pipeline_parallel"], - refit_urls=["http://receiver"], ) generation.worker_group.run_all_workers_single_data.assert_any_call( "configure_zmq_sparse_refit_relay", From 0ce04977ec5c65ee21bdc7aea2f768f36c780d14 Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Thu, 16 Jul 2026 12:21:19 -0700 Subject: [PATCH 14/18] Add nightly tests Signed-off-by: Hollow Man --- ...30ba3b-4n8g-megatron-zmq-noncolocated.yaml | 38 ++++++++++++++++ ...3-30ba3b-4n8g-megatron-zmq-noncolocated.sh | 45 +++++++++++++++++++ tests/test_suites/nightly.txt | 1 + 3 files changed, 84 insertions(+) create mode 100644 examples/configs/recipes/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.yaml create mode 100755 tests/test_suites/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.sh diff --git a/examples/configs/recipes/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.yaml b/examples/configs/recipes/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.yaml new file mode 100644 index 00000000000..a8c1fadd74e --- /dev/null +++ b/examples/configs/recipes/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.yaml @@ -0,0 +1,38 @@ +defaults: ./performance/grpo-qwen3-30ba3b-4n8g.yaml + +grpo: + num_prompts_per_step: 16 + num_generations_per_prompt: 8 + max_num_steps: 50 + val_period: 1000 + +checkpointing: + checkpoint_dir: results/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated + +policy: + train_global_batch_size: 128 + generation_batch_size: 16 + max_total_sequence_length: 2048 + sequence_packing: + train_mb_tokens: 2048 + logprob_mb_tokens: 4096 + megatron_cfg: + activation_checkpointing: true + env_vars: + PYTORCH_CUDA_ALLOC_CONF: expandable_segments:True + generation: + refit_transport: vllm_zmq_sparse + delta_compression: + encoding: xor + sparse_bucket_size_bytes: 536870912 + colocated: + enabled: false + resources: + gpus_per_node: 8 + num_nodes: 2 + +logger: + log_dir: logs/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated + wandb: + project: nemo-rl-refit + name: grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated diff --git a/tests/test_suites/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.sh b/tests/test_suites/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.sh new file mode 100755 index 00000000000..2bea0a11ace --- /dev/null +++ b/tests/test_suites/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.sh @@ -0,0 +1,45 @@ +#!/bin/bash +SCRIPT_DIR=$( cd -- "$( dirname -- "${BASH_SOURCE[0]}" )" &> /dev/null && pwd) +source $SCRIPT_DIR/common.env + +# ===== BEGIN CONFIG ===== +NUM_NODES=4 +STEPS_PER_RUN=50 +MAX_STEPS=50 +NUM_RUNS=$(( (MAX_STEPS + STEPS_PER_RUN - 1) / STEPS_PER_RUN )) +NUM_MINUTES=65 +# ===== END CONFIG ===== + +exit_if_max_steps_reached + +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-refit \ + logger.wandb.name=$EXP_NAME \ + logger.monitor_gpus=True \ + logger.tensorboard_enabled=True \ + checkpointing.enabled=False \ + $@ \ + 2>&1 | tee $RUN_LOG + +uv run tests/json_dump_tb_logs.py $LOG_DIR --output_path $JSON_METRICS + +MAX_RECORDED_STEP=$(jq -r 'if has("train/loss") then (."train/loss" | keys | map(tonumber) | max // 0) else 0 end' $JSON_METRICS) +if [[ $MAX_RECORDED_STEP -lt $MAX_STEPS ]]; then + echo "[ERROR] Expected train/loss through step $MAX_STEPS, found step $MAX_RECORDED_STEP" + exit 1 +fi + +uv run tests/check_metrics.py $JSON_METRICS \ + 'median(data["train/token_mult_prob_error"]) < 1.03' \ + "data[\"train/token_mult_prob_error\"][\"$MAX_STEPS\"] < 1.03" \ + 'ratio_above(data["train/token_mult_prob_error"], 1.03) < 0.05' \ + "data[\"train/reward\"][\"$MAX_STEPS\"] > 0.2" \ + 'min(data["refit/transfer/payloads"]) > 0' \ + 'min(data["refit/transfer/relay_flush_s"]) > 0' \ + 'min(data["refit/delta/changed_pct"]) > 0' \ + 'max(data["refit/delta/changed_pct"]) <= 5' diff --git a/tests/test_suites/nightly.txt b/tests/test_suites/nightly.txt index 81e8ebc90c4..95db3bfbf34 100644 --- a/tests/test_suites/nightly.txt +++ b/tests/test_suites/nightly.txt @@ -93,6 +93,7 @@ tests/test_suites/llm/grpo-qwen3-8b-base-dapo-2n8g-long-megatron-qa-nvfp4-w4a16. # Non-colocated tests/test_suites/llm/grpo-llama3.1-8b-instruct-2n8g-fsdp2tp1-noncolocated.sh tests/test_suites/llm/grpo-llama3.2-1b-instruct-2n8g-megatron_generation-noncolocated.sh +tests/test_suites/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.sh # Nemotron Super 49B #https://github.com/NVIDIA-NeMo/RL/issues/1374 From 011a259f30f546133885e767479372984681ffb9 Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Thu, 16 Jul 2026 19:15:03 -0700 Subject: [PATCH 15/18] Increase nightly compute Signed-off-by: Hollow Man --- tests/unit/test_recipes_and_test_suites.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tests/unit/test_recipes_and_test_suites.py b/tests/unit/test_recipes_and_test_suites.py index 3e7f01365cc..968a90265a1 100644 --- a/tests/unit/test_recipes_and_test_suites.py +++ b/tests/unit/test_recipes_and_test_suites.py @@ -255,7 +255,7 @@ def test_all_recipe_yamls_accounted_for_in_test_suites( ) -def test_nightly_compute_stays_below_3390_hours(nightly_test_suite, tracker): +def test_nightly_compute_stays_below_3420_hours(nightly_test_suite, tracker): command = f"DRYRUN=1 HF_HOME=... HF_DATASETS_CACHE=... CONTAINER= ACCOUNT= PARTITION= ./tools/launch {' '.join(nightly_test_suite)}" print(f"Running command: {command}") @@ -287,8 +287,8 @@ def test_nightly_compute_stays_below_3390_hours(nightly_test_suite, tracker): f"Last line of output was not as expected: '{last_line}'" ) total_gpu_hours = float(last_line.split(":")[-1].strip()) - assert total_gpu_hours <= 3390, ( - f"Total GPU hours exceeded 3390: {last_line}. We should revisit the test suites to reduce the total GPU hours." + assert total_gpu_hours <= 3420, ( + f"Total GPU hours exceeded 3420: {last_line}. We should revisit the test suites to reduce the total GPU hours." ) tracker.track("total_nightly_gpu_hours", total_gpu_hours) From 5a257c434c0f50dde445e54722f7a68f26ad0be0 Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Fri, 17 Jul 2026 13:35:49 -0700 Subject: [PATCH 16/18] Clean up test command Signed-off-by: Hollow Man --- docs/design-docs/sparse-delta-refit.md | 20 -------------------- 1 file changed, 20 deletions(-) diff --git a/docs/design-docs/sparse-delta-refit.md b/docs/design-docs/sparse-delta-refit.md index 15642f19b9d..d3877049082 100644 --- a/docs/design-docs/sparse-delta-refit.md +++ b/docs/design-docs/sparse-delta-refit.md @@ -386,26 +386,6 @@ deterministic chunk ownership, transactional baseline updates, transport cleanup on success and failure, and unchanged producer/receiver overlap. Do not modify Megatron Bridge for a transport-specific hook. -Run the focused suite: - -```bash -uv run --extra vllm pytest -q \ - tests/unit/utils/test_weight_transfer_stream.py \ - tests/unit/models/policy/test_megatron_remote_sparse_refit.py \ - tests/unit/models/generation/test_vllm_sparse_refit.py \ - tests/unit/weight_sync/test_vllm_remote_sparse_weight_synchronizer.py - -uv run --extra vllm pytest -q -m vllm \ - tests/unit/models/generation/test_vllm_sparse_delta.py - -uv run ruff check \ - nemo_rl/utils/weight_transfer_{http,sparse_codec,stream,zmq}.py \ - nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py \ - nemo_rl/models/generation/vllm/vllm_{sparse_refit,sparse_delta}.py \ - nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py \ - tools/refit_bandwidth_calculator.py -``` - On the target topology, verify the exact commit, image digest, and checkpoint revision; validate fresh starts and same-version resumes; compare two balanced repetitions with an equivalent NCCL or full control; and require the requested From caeee9076abc3cfae1a161982d719b26f2bce93c Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Fri, 17 Jul 2026 16:39:46 -0700 Subject: [PATCH 17/18] Address review issues Signed-off-by: Hollow Man --- docs/design-docs/sparse-delta-refit.md | 142 ++++++++++---- examples/configs/grpo_math_1B.yaml | 2 +- ...egatron-zmq-deltaweight-noncolocated.yaml} | 15 +- nemo_rl/algorithms/distillation.py | 8 + nemo_rl/algorithms/grpo.py | 6 + nemo_rl/algorithms/ppo.py | 8 + nemo_rl/models/generation/vllm/config.py | 59 +++++- .../generation/vllm/vllm_sparse_refit.py | 56 ++++-- .../policy/workers/megatron_policy_worker.py | 6 +- .../workers/megatron_remote_sparse_refit.py | 13 +- nemo_rl/utils/weight_transfer_sparse_codec.py | 25 ++- nemo_rl/utils/weight_transfer_stream.py | 63 +++--- nemo_rl/utils/weight_transfer_zmq.py | 32 ++-- .../vllm_remote_sparse_weight_synchronizer.py | 28 ++- ...-megatron-zmq-deltaweight-noncolocated.sh} | 0 tests/test_suites/nightly.txt | 2 +- .../generation/test_vllm_sparse_refit.py | 5 +- .../test_megatron_remote_sparse_refit.py | 28 ++- .../unit/reference_configs/grpo_math_1B.yaml | 2 +- .../unit/utils/test_weight_transfer_stream.py | 180 ++++++++++++------ ..._vllm_remote_sparse_weight_synchronizer.py | 17 +- 21 files changed, 482 insertions(+), 215 deletions(-) rename examples/configs/recipes/llm/{grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.yaml => grpo-qwen3-30ba3b-4n8g-megatron-zmq-deltaweight-noncolocated.yaml} (64%) rename tests/test_suites/llm/{grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.sh => grpo-qwen3-30ba3b-4n8g-megatron-zmq-deltaweight-noncolocated.sh} (100%) diff --git a/docs/design-docs/sparse-delta-refit.md b/docs/design-docs/sparse-delta-refit.md index d3877049082..82f823bab74 100644 --- a/docs/design-docs/sparse-delta-refit.md +++ b/docs/design-docs/sparse-delta-refit.md @@ -18,7 +18,7 @@ Remote sparse refit requires: - the same initial HF checkpoint on both clusters; - BF16 or FP16 unquantized rollout weights; - `kv_cache_dtype: auto`; and -- a `delta_compression` configuration. +- `refit_transport: vllm_s3_sparse` or `vllm_zmq_sparse`. Validation rejects `quant_cfg`, `real_quant`, colocated or non-Megatron deployments, FP8 rollout weights, and FP8 KV-cache scales. The codec and @@ -28,34 +28,55 @@ Synchronous and asynchronous vLLM engines are supported, but the weight-version transition remains synchronous: generation pauses until every payload is applied and the global flush completes. +The coordinator is currently integrated only with GRPO. PPO and distillation +reject `refit_transport` during setup; extending them is tracked in +[#3275](https://github.com/NVIDIA-NeMo/RL/issues/3275). + ## Architecture ```mermaid flowchart LR + SYNC["GRPO weight synchronizer"] + subgraph P["Megatron policy cluster"] - B["Full canonical Bridge export"] - C["Chunk-sharded canonical HF baseline"] - E["Compare, encode, and compress"] - B --> C - C --> E + MB["All policy ranks
Megatron Bridge HF export"] + OWN["Deterministic chunk owner
chunk_index modulo producers"] + BASE["Distributed canonical
CPU or mmap baseline"] + PIPE["Shared bounded pipeline
compare, encode, zstd"] + MB -->|canonical HF chunks| OWN + OWN -->|owned chunks| BASE + BASE -->|changed locations and bits| PIPE end - S["S3 object and HTTP manifest"] - Z["ZeroMQ relay"] + subgraph T["Value plane"] + S3["S3 objects
one upload per owned payload"] + Z0["ZeroMQ root relay
one cross-cluster send"] + ZT["Binary relay tree"] + Z0 --> ZT + end subgraph G["vLLM generation cluster"] - H["HTTP receiver"] - Q["Compact node staging and bounded FIFO apply queue"] - A["Reusable dense scratch and native load_weights()"] - H --> Q --> A + HTTP["HTTP control plane
prepare and global flush"] + RX["One receiver per inference node"] + STAGE["Compact payload staging
and bounded FIFO queue"] + RPC["Node-local collective RPC"] + APPLY["Each vLLM rank
reusable scratch + native load_weights()"] + HTTP -.-> RX + RX --> STAGE --> RPC --> APPLY end - E -->|S3| S --> H - E -->|ZeroMQ| Z --> H + SYNC -. "prepare, flush, commit" .-> HTTP + SYNC -. "start stream" .-> PIPE + PIPE -->|compressed body| S3 + S3 -->|HTTP manifest, then GET| RX + PIPE -->|multipart DATA frame| Z0 + ZT -->|compressed body| RX + APPLY -. "success metrics" .-> SYNC + SYNC -. "commit exact source bits" .-> BASE ``` -*Figure 1. S3 and ZeroMQ share the exporter, codec, receiver, native-loader -apply engine, and commit protocol.* +*Figure 1. Control traffic is dashed. Payload data is solid. S3 and ZeroMQ +share every component except the value plane.* | Responsibility | Implementation | |---|---| @@ -99,7 +120,7 @@ Mamba, padded or tied weights, grouped exports, adapters, and custom Bridge postprocessing all follow Bridge's canonical export semantics. Baselines use file-backed `torch.from_file` tensors by default; -`NRL_REFIT_BASELINE_IN_MEMORY=1` keeps them in RAM. File backing reduces +`refit_cfg.baseline.in_memory: true` keeps them in RAM. File backing reduces anonymous resident-memory pressure but does not reduce logical baseline bytes. Baseline initialization also returns each canonical tensor's name, shape, and @@ -236,6 +257,9 @@ commit exact pending baseline bits in background CPU threads. > receiver accepts any payload, reload that receiver from a known-good weight > version before retrying. This is mandatory for XOR, because replaying an > already-applied XOR reverts those bits. Replaying overwrite is safe. +> The synchronizer records this state as poisoned and rejects subsequent syncs +> with recovery instructions. In-place recovery is tracked in +> [#3274](https://github.com/NVIDIA-NeMo/RL/issues/3274). ## Payload and native apply @@ -245,6 +269,32 @@ Each serialized payload is: (packed_location_bytes, packed_value_groups, tensor_metadata) ``` +```mermaid +flowchart TB + subgraph S3F["S3 transport"] + SO["Object body
zstd-compressed torch payload"] + SM["HTTP manifest
bucket, region, key, checksum,
verification_candidates"] + end + + subgraph ZF["ZeroMQ DATA multipart"] + ZK["Frame 1: DATA"] + ZM["Frame 2: JSON metadata
transfer, producer, payload IDs,
checksum, samples, optional API key"] + ZB["Frame 3: zstd-compressed torch payload"] + end + + SO --> DECOMP["Checksum then zstd decode"] + SM -. "locates and authenticates object" .-> SO + ZK --> ZM --> ZB --> DECOMP + DECOMP --> PT["torch payload tuple"] + PT --> PI["packed location bytes
range or uint16/32/64 deltas"] + PT --> PV["value groups by dtype
XOR or overwrite bits"] + PT --> PM["per-tensor metadata
HF name, shape, offsets, operation"] +``` + +*Figure 2. S3 separates the object body from its control-plane manifest; +ZeroMQ carries equivalent identity and body data as multipart frames. Both +decode into the same codec tuple.* + Contiguous locations use a range encoding. Other sorted locations are delta-encoded into the smallest lossless unsigned width among 16, 32, and 64 bits. Metadata carries the HF name and shape, dtype, value offsets, location @@ -285,9 +335,15 @@ policy: generation: backend: vllm refit_transport: vllm_s3_sparse # or vllm_zmq_sparse - delta_compression: - encoding: xor # overwrite is selected automatically for opaque loaders - sparse_bucket_size_bytes: 536870912 + refit_cfg: + delta_compression: + encoding: xor # overwrite is selected automatically for opaque loaders + storage: + s3_bucket: my-refit-bucket # required only for vllm_s3_sparse + s3_region: us-east-1 + baseline: + in_memory: false + verify_samples_per_payload: 0 colocated: enabled: false vllm_cfg: @@ -299,32 +355,44 @@ policy: zmq_refit_server_port: null ``` -S3 requires `NRL_REFIT_S3_BUCKET`; region and key prefix default to `us-east-1` -and `nemo-rl-refit`. ZeroMQ requires routable TCP access to the relay port. The -HTTP and ZeroMQ servers are plaintext, so use a trusted or encrypted network. -When `http_refit_api_key_env_var` is set, the named variable must contain the -same nonempty token on producers and receivers. +`refit_cfg` is optional; its Pydantic models resolve and log all defaults. S3 +fails during setup unless `refit_cfg.storage.s3_bucket` is nonempty. Its region +and key prefix default to `us-east-1` and `nemo-rl-refit`. ZeroMQ requires +routable TCP access to the relay port. The HTTP and ZeroMQ servers are +plaintext, so use a trusted or encrypted network. When +`http_refit_api_key_env_var` is set, the named variable must contain the same +nonempty token on producers and receivers. Binding either server to all +interfaces without a key emits a warning. | Control | Default | |---|---:| -| `delta_compression.sparse_bucket_size_bytes` | 512 MiB | -| `NRL_REFIT_S3_EXPORT_CHUNK_BYTES` | 64 MiB | -| `NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES` | 256 MiB | -| `NRL_REFIT_{S3,ZMQ}_ENCODE_WORKERS` | 2-8 from CPU count | -| `NRL_REFIT_S3_UPLOAD_WORKERS` | 4-32 from CPU count | -| `NRL_REFIT_ZMQ_SEND_WORKERS` | 4 | -| `NRL_REFIT_ZMQ_RELAY_PAYLOAD_WORKERS` | 16 | -| `NRL_REFIT_ZMQ_RELAY_FORWARD_WORKERS` | 8 | -| `NRL_REFIT_APPLY_QUEUE_DEPTH` / `NRL_REFIT_APPLY_BATCH_SIZE` | 32 / 8 | -| `NRL_REFIT_PARTITION_WORKERS` | 2-8 from CPU count | -| `NRL_REFIT_{S3,ZMQ}_ZSTD_THREADS` | 0 | -| `NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD` | 0 | +| `refit_cfg.delta_compression.encoding` | `xor` | +| `refit_cfg.storage.s3_bucket` | none; required for S3 | +| `refit_cfg.storage.s3_region` | `us-east-1` | +| `refit_cfg.baseline.in_memory` | `false` | +| `refit_cfg.verify_samples_per_payload` | `0` | Export chunks are capped by `sparse_bucket_size_bytes` and the packed tensor limit. The S3 defaults were selected by balanced 120B sweeps. Increase one concurrency control at a time; excess parallelism moves the bottleneck into host memory, Bridge export, relay-tree forwarding, or receiver apply. +Run the checked-in ZeroMQ recipe with: + +```bash +uv run python examples/run_grpo.py \ + --config examples/configs/recipes/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-deltaweight-noncolocated.yaml +``` + +For a diagnostic run, enable bounded transmitted-delta sampling through the +logged config: + +```bash +uv run python examples/run_grpo.py \ + --config examples/configs/recipes/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-deltaweight-noncolocated.yaml \ + policy.generation.refit_cfg.verify_samples_per_payload=32 +``` + ## Metrics and profiling | Signal | Meaning | diff --git a/examples/configs/grpo_math_1B.yaml b/examples/configs/grpo_math_1B.yaml index b901812bd06..44ae519a3f9 100644 --- a/examples/configs/grpo_math_1B.yaml +++ b/examples/configs/grpo_math_1B.yaml @@ -340,7 +340,7 @@ policy: stop_token_ids: null stop_strings: null refit_transport: null # Set to "vllm_s3_sparse" or "vllm_zmq_sparse" for remote sparse-delta refit. - delta_compression: null # Set {} for XOR and the 512 MiB bucket; null uses the existing refit path. + refit_cfg: null # Optional tuning and storage settings for remote sparse-delta refit. mcore_generation_config: async_engine: false max_model_len: ${policy.max_total_sequence_length} # Engine-side max sequence length. diff --git a/examples/configs/recipes/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.yaml b/examples/configs/recipes/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-deltaweight-noncolocated.yaml similarity index 64% rename from examples/configs/recipes/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.yaml rename to examples/configs/recipes/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-deltaweight-noncolocated.yaml index a8c1fadd74e..18ff88ca112 100644 --- a/examples/configs/recipes/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.yaml +++ b/examples/configs/recipes/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-deltaweight-noncolocated.yaml @@ -7,7 +7,7 @@ grpo: val_period: 1000 checkpointing: - checkpoint_dir: results/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated + checkpoint_dir: results/grpo-qwen3-30ba3b-4n8g-megatron-zmq-deltaweight-noncolocated policy: train_global_batch_size: 128 @@ -22,9 +22,12 @@ policy: PYTORCH_CUDA_ALLOC_CONF: expandable_segments:True generation: refit_transport: vllm_zmq_sparse - delta_compression: - encoding: xor - sparse_bucket_size_bytes: 536870912 + refit_cfg: + delta_compression: + encoding: xor + verify_samples_per_payload: 0 + baseline: + in_memory: false colocated: enabled: false resources: @@ -32,7 +35,7 @@ policy: num_nodes: 2 logger: - log_dir: logs/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated + log_dir: logs/grpo-qwen3-30ba3b-4n8g-megatron-zmq-deltaweight-noncolocated wandb: project: nemo-rl-refit - name: grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated + name: grpo-qwen3-30ba3b-4n8g-megatron-zmq-deltaweight-noncolocated diff --git a/nemo_rl/algorithms/distillation.py b/nemo_rl/algorithms/distillation.py index 786ffb4548c..7818a8dd418 100644 --- a/nemo_rl/algorithms/distillation.py +++ b/nemo_rl/algorithms/distillation.py @@ -217,6 +217,14 @@ def setup( assert generation_config is not None, ( "A generation config in the PolicyConfig is required for distillation" ) + if ( + generation_config["backend"] == "vllm" + and cast(VllmConfig, generation_config).get("refit_transport") is not None + ): + raise ValueError( + "Remote sparse refit is currently supported only by GRPO; distillation " + "support is tracked in https://github.com/NVIDIA-NeMo/RL/issues/3275." + ) # Disallow SP + packing for dtensor path for cfg, who in ((policy_config, "student"), (teacher_config, "teacher")): diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index 4f7e0b58aa9..2e3a472809f 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -101,6 +101,7 @@ from nemo_rl.models.generation.sglang.config import SGLangConfig from nemo_rl.models.generation.sglang.sglang_generation import SGLangGeneration from nemo_rl.models.generation.vllm import VllmConfig, VllmGeneration +from nemo_rl.models.generation.vllm.config import normalize_vllm_refit_config from nemo_rl.models.megatron.router_replay import ( configure_vllm_for_router_replay, router_replay_enabled, @@ -354,6 +355,8 @@ def setup( assert generation_config is not None, ( "A generation config in the PolicyConfig is required for GRPO" ) + if generation_config["backend"] == "vllm": + normalize_vllm_refit_config(cast(VllmConfig, generation_config)) # Set seed for all random number generators set_seed(grpo_config["seed"]) @@ -1298,6 +1301,8 @@ def init_vllm_then_policy(): t0 = time.perf_counter() assert isinstance(policy_generation, VllmGeneration) assert remote_synchronizer_cls is not None + refit_config = generation_config["refit_cfg"] + assert refit_config is not None policy_generation.weight_synchronizer = remote_synchronizer_cls( policy, policy_generation, @@ -1305,6 +1310,7 @@ def init_vllm_then_policy(): api_key_env_var=generation_config["vllm_cfg"].get( "http_refit_api_key_env_var" ), + request_timeout_s=refit_config.request_timeout_s, baseline_init_refs=remote_baseline_init_refs, ) policy_generation.weight_synchronizer.init_communicator() diff --git a/nemo_rl/algorithms/ppo.py b/nemo_rl/algorithms/ppo.py index 80340bbbf30..5bef74b5cef 100644 --- a/nemo_rl/algorithms/ppo.py +++ b/nemo_rl/algorithms/ppo.py @@ -228,6 +228,14 @@ def setup( assert generation_config is not None, ( "A generation config in the PolicyConfig is required for PPO" ) + if ( + generation_config["backend"] == "vllm" + and cast(VllmConfig, generation_config).get("refit_transport") is not None + ): + raise ValueError( + "Remote sparse refit is currently supported only by GRPO; PPO support " + "is tracked in https://github.com/NVIDIA-NeMo/RL/issues/3275." + ) if "megatron_cfg" in policy_config and policy_config["megatron_cfg"]["enabled"]: policy_megatron_config = cast(MegatronConfig, policy_config["megatron_cfg"]) diff --git a/nemo_rl/models/generation/vllm/config.py b/nemo_rl/models/generation/vllm/config.py index 56ac513176b..6b831cce923 100644 --- a/nemo_rl/models/generation/vllm/config.py +++ b/nemo_rl/models/generation/vllm/config.py @@ -14,10 +14,12 @@ from typing import Any, Literal, NotRequired, TypedDict -from pydantic import BaseModel, PositiveInt +from pydantic import BaseModel, Field, NonNegativeInt, PositiveFloat, PositiveInt from nemo_rl.models.generation.interfaces import GenerationConfig +VllmRefitTransportName = Literal["s3", "zmq"] + class VllmSpecificArgs(TypedDict): tensor_parallel_size: int @@ -63,6 +65,50 @@ class VllmSpecificArgs(TypedDict): class VllmDeltaCompressionConfig(BaseModel, extra="allow"): encoding: Literal["xor", "overwrite"] = "xor" sparse_bucket_size_bytes: PositiveInt = 512 * 1024**2 + export_chunk_bytes: dict[str, PositiveInt] = Field( + default_factory=lambda: {"s3": 64 * 1024**2, "zmq": 256 * 1024**2} + ) + zstd_threads: dict[str, NonNegativeInt] = Field( + default_factory=lambda: {"s3": 0, "zmq": 0} + ) + + +class VllmRefitStorageConfig(BaseModel, extra="allow"): + s3_bucket: str | None = None + s3_region: str = "us-east-1" + s3_prefix: str = "nemo-rl-refit" + staging_dir: str = "/dev/shm" + + +class VllmRefitBaselineConfig(BaseModel, extra="allow"): + in_memory: bool = False + mmap_dir: str | None = None + + +class VllmRefitTuningConfig(BaseModel, extra="allow"): + encode_workers: dict[str, PositiveInt] = Field( + default_factory=lambda: {"s3": 8, "zmq": 8} + ) + transfer_workers: dict[str, PositiveInt] = Field( + default_factory=lambda: {"s3": 32, "zmq": 4} + ) + zmq_retries: NonNegativeInt = 3 + zmq_relay_payload_workers: PositiveInt = 16 + zmq_relay_forward_workers: PositiveInt = 8 + apply_queue_depth: PositiveInt = 32 + apply_batch_size: PositiveInt = 8 + partition_workers: PositiveInt = 8 + + +class VllmRefitConfig(BaseModel, extra="allow"): + delta_compression: VllmDeltaCompressionConfig = Field( + default_factory=VllmDeltaCompressionConfig + ) + storage: VllmRefitStorageConfig = Field(default_factory=VllmRefitStorageConfig) + baseline: VllmRefitBaselineConfig = Field(default_factory=VllmRefitBaselineConfig) + tuning: VllmRefitTuningConfig = Field(default_factory=VllmRefitTuningConfig) + verify_samples_per_payload: NonNegativeInt = 0 + request_timeout_s: PositiveFloat = 600.0 class VllmConfig(GenerationConfig): @@ -70,7 +116,7 @@ class VllmConfig(GenerationConfig): vllm_kwargs: NotRequired[dict[str, Any]] # Null uses NCCL; remote sparse refit supports S3 or ZeroMQ value planes. refit_transport: NotRequired[Literal["vllm_s3_sparse", "vllm_zmq_sparse"] | None] - delta_compression: NotRequired[VllmDeltaCompressionConfig | None] + refit_cfg: NotRequired[VllmRefitConfig | None] # quantization config quant_cfg: NotRequired[str | None] @@ -79,3 +125,12 @@ class VllmConfig(GenerationConfig): # modules. This is intended for ModelOpt NVFP4 rollout experiments. real_quant: NotRequired[bool] real_quant_ignore: NotRequired[list[str]] + + +def normalize_vllm_refit_config(config: VllmConfig) -> VllmRefitConfig | None: + """Resolve sparse-refit defaults into the generation config.""" + if config.get("refit_transport") is None: + return None + refit_config = VllmRefitConfig.model_validate(config.get("refit_cfg") or {}) + config["refit_cfg"] = refit_config + return refit_config diff --git a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py index 9965007eda9..b298d0674d1 100644 --- a/nemo_rl/models/generation/vllm/vllm_sparse_refit.py +++ b/nemo_rl/models/generation/vllm/vllm_sparse_refit.py @@ -15,12 +15,15 @@ """Remote sparse-refit receiver lifecycle for vLLM generation workers.""" import asyncio +import hmac +import logging import os import tempfile import threading import time from collections.abc import Mapping from concurrent.futures import Future, ThreadPoolExecutor +from functools import cache from typing import Any, Literal, NamedTuple, cast import uvicorn @@ -33,6 +36,7 @@ _get_free_port_local, _get_node_ip_local, ) +from nemo_rl.models.generation.vllm.config import VllmRefitConfig from nemo_rl.utils import weight_transfer_sparse_codec as sparse_codec from nemo_rl.utils.weight_transfer_http import ( G_VLLM_REFIT_API_KEY_HEADER, @@ -46,12 +50,22 @@ from nemo_rl.utils.weight_transfer_stream import ( decode_sparse_payload, download_s3_refit_payload, - refit_env_int, ) from nemo_rl.utils.weight_transfer_zmq import ( ZmqSparseRefitServer, ) +logger = logging.getLogger(__name__) + + +@cache +def _warn_unauthenticated_refit_server(transport: str) -> None: + logger.warning( + "%s sparse-refit server is binding 0.0.0.0 without an API key; " + "weight-write endpoints are reachable from the network.", + transport, + ) + class _StagedSparsePayload(NamedTuple): path: str @@ -90,6 +104,10 @@ class VllmSparseRefitReceiver: def __init__(self, worker: Any) -> None: self._worker = worker + self._refit_config = VllmRefitConfig.model_validate( + worker.cfg.get("refit_cfg") or {} + ) + tuning = self._refit_config.tuning self._refit_apply_queue_condition = threading.Condition() self._refit_apply_executor = ThreadPoolExecutor( max_workers=1, @@ -101,23 +119,14 @@ def __init__(self, worker: Any) -> None: ] = [] self._refit_seen_payloads: dict[tuple[str, int, int], str] = {} self._refit_workers_share_node = False - self._refit_apply_queue_depth = refit_env_int( - "NRL_REFIT_APPLY_QUEUE_DEPTH", default=32 - ) - self._refit_apply_batch_size = refit_env_int( - "NRL_REFIT_APPLY_BATCH_SIZE", default=8 - ) + self._refit_apply_queue_depth = tuning.apply_queue_depth + self._refit_apply_batch_size = tuning.apply_batch_size self._refit_partition_executor = ThreadPoolExecutor( - max_workers=refit_env_int( - "NRL_REFIT_PARTITION_WORKERS", - default=max(2, min(8, os.cpu_count() or 8)), - ), + max_workers=tuning.partition_workers, thread_name_prefix="nrl-vllm-sparse-partition", ) self._refit_verification_candidates = 0 - self._refit_batch_staging_dir = ( - os.getenv("NRL_REFIT_BATCH_STAGING_DIR") or "/dev/shm" - ) + self._refit_batch_staging_dir = self._refit_config.storage.staging_dir self._refit_http_server: tuple[Any, threading.Thread, str] | None = None self._zmq_refit_server: tuple[ZmqSparseRefitServer, str] | None = None self._refit_async_loop: asyncio.AbstractEventLoop | None = None @@ -419,9 +428,9 @@ async def respond( ) -> JSONResponse: if cfg["vllm_cfg"]["async_engine"]: self._refit_async_loop = asyncio.get_running_loop() - if ( - token is not None - and raw_request.headers.get(G_VLLM_REFIT_API_KEY_HEADER) != token + supplied_token = raw_request.headers.get(G_VLLM_REFIT_API_KEY_HEADER) + if token is not None and ( + supplied_token is None or not hmac.compare_digest(token, supplied_token) ): return JSONResponse( content={"ok": False, "error": "unauthorized"}, status_code=403 @@ -486,8 +495,14 @@ def start_zmq_sparse_refit_relay(self) -> str: self._apply_zmq_payload, bind_address=f"tcp://0.0.0.0:{port}", api_key_env_var=cfg["vllm_cfg"].get("http_refit_api_key_env_var"), - timeout_s=float(os.getenv("NRL_REFIT_ZMQ_TIMEOUT_S") or 600.0), + timeout_s=self._refit_config.request_timeout_s, + tuning=self._refit_config.tuning, ) + if ( + vllm_refit_api_key(cfg["vllm_cfg"].get("http_refit_api_key_env_var")) + is None + ): + _warn_unauthenticated_refit_server("ZeroMQ") server.start() address = f"tcp://{_get_node_ip_local()}:{port}" self._zmq_refit_server = (server, address) @@ -523,6 +538,11 @@ def _setup_vllm_refit_server(self) -> None: cfg.get("port_range_low", DEFAULT_GENERATION_PORT_RANGE_LOW), cfg.get("port_range_high", DEFAULT_GENERATION_PORT_RANGE_HIGH), ) + if ( + vllm_refit_api_key(cfg["vllm_cfg"].get("http_refit_api_key_env_var")) + is None + ): + _warn_unauthenticated_refit_server("HTTP") server = uvicorn.Server( uvicorn.Config( app, diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 9f775e1db32..dea0c079a14 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -1842,9 +1842,9 @@ def _require_remote_sparse_refit(self) -> Any: MegatronRemoteSparseRefit, ) - self._remote_sparse_refit = MegatronRemoteSparseRefit( - self, self.cfg["generation"]["delta_compression"] - ) + refit_config = self.cfg["generation"]["refit_cfg"] + assert refit_config is not None + self._remote_sparse_refit = MegatronRemoteSparseRefit(self, refit_config) return self._remote_sparse_refit def finish_remote_sparse_delta_sync(self, *, succeeded: bool) -> None: diff --git a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py index 2b86eba287d..dced0511ee0 100644 --- a/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py +++ b/nemo_rl/models/policy/workers/megatron_remote_sparse_refit.py @@ -18,7 +18,10 @@ import torch -from nemo_rl.models.generation.vllm.config import VllmDeltaCompressionConfig +from nemo_rl.models.generation.vllm.config import ( + VllmRefitConfig, + VllmRefitTransportName, +) from nemo_rl.utils.weight_transfer_sparse_codec import DeltaCompressionTracker from nemo_rl.utils.weight_transfer_stream import ( init_sparse_delta_baseline_from_iterator, @@ -28,16 +31,16 @@ class MegatronRemoteSparseRefit: - def __init__(self, worker: Any, delta_config: VllmDeltaCompressionConfig) -> None: + def __init__(self, worker: Any, refit_config: VllmRefitConfig) -> None: self._worker = worker - self._tracker = DeltaCompressionTracker(delta_config.model_dump()) + self._tracker = DeltaCompressionTracker(refit_config) def initialize_baseline( self, *, shard_rank: int, shard_count: int, - transport: str, + transport: VllmRefitTransportName, ) -> dict[str, tuple[tuple[int, ...], torch.dtype]]: init_sparse_delta_baseline_from_iterator( self._worker._iter_params_with_optional_kv_scales(), @@ -53,7 +56,7 @@ def initialize_baseline( def stream( self, - transport: str, + transport: VllmRefitTransportName, targets: list[str], *, transfer_id: str, diff --git a/nemo_rl/utils/weight_transfer_sparse_codec.py b/nemo_rl/utils/weight_transfer_sparse_codec.py index 4a9d5ecda3b..8e8eae54735 100644 --- a/nemo_rl/utils/weight_transfer_sparse_codec.py +++ b/nemo_rl/utils/weight_transfer_sparse_codec.py @@ -12,16 +12,17 @@ # See the License for the specific language governing permissions and # limitations under the License. -import os import tempfile import threading -from collections.abc import Iterable, Mapping +from collections.abc import Iterable from concurrent.futures import ThreadPoolExecutor from typing import Any, Literal import numpy as np import torch +from nemo_rl.models.generation.vllm.config import VllmRefitConfig + NamedTensor = tuple[str, torch.Tensor] TensorBatch = list[NamedTensor] SparseOperation = Literal["xor", "overwrite"] @@ -208,20 +209,16 @@ class DeltaCompressionTracker: def __init__( self, - config: Mapping[str, Any], + config: VllmRefitConfig, ) -> None: - self.sparse_bucket_size_bytes = int(config["sparse_bucket_size_bytes"]) - if self.sparse_bucket_size_bytes < 1: - raise ValueError("delta_compression.sparse_bucket_size_bytes must be >= 1") - self.encoding = sparse_operation(config["encoding"]) + self.refit_config = config + delta_config = config.delta_compression + self.sparse_bucket_size_bytes = delta_config.sparse_bucket_size_bytes + self.encoding = sparse_operation(delta_config.encoding) self.overwrite_names: frozenset[str] = frozenset() - self.verification_samples = int( - os.getenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "0") - ) - if self.verification_samples < 0: - raise ValueError("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD must be >= 0") - self.baseline_in_memory = os.getenv("NRL_REFIT_BASELINE_IN_MEMORY") == "1" - self.baseline_mmap_dir = os.getenv("NRL_REFIT_BASELINE_MMAP_DIR") + self.verification_samples = config.verify_samples_per_payload + self.baseline_in_memory = config.baseline.in_memory + self.baseline_mmap_dir = config.baseline.mmap_dir self.baseline: dict[str, torch.Tensor] = {} self._pending_updates: dict[str, tuple[torch.Tensor, torch.Tensor]] = {} self._pending_updates_lock = threading.Lock() diff --git a/nemo_rl/utils/weight_transfer_stream.py b/nemo_rl/utils/weight_transfer_stream.py index 2645f884454..b39d0804df9 100644 --- a/nemo_rl/utils/weight_transfer_stream.py +++ b/nemo_rl/utils/weight_transfer_stream.py @@ -16,7 +16,6 @@ import hashlib import io -import os import threading import time from collections.abc import Iterable, Iterator, Mapping, Sequence @@ -30,6 +29,7 @@ import torch import zstandard +from nemo_rl.models.generation.vllm.config import VllmRefitTransportName from nemo_rl.utils.packed_tensor import get_target_packed_tensor_size from nemo_rl.utils.weight_transfer_http import ( G_VLLM_REFIT_S3_MANIFEST_PATH, @@ -47,6 +47,8 @@ ) _STREAM_LOCAL = threading.local() +# A 64 MiB CRT part and 2 GiB client cap allow 32 in-flight parts, matching the +# upload concurrency selected by the balanced 120B S3 sweeps while bounding RAM. _S3_PART_SIZE = 64 * 1024**2 _S3_MEMORY_LIMIT = 2 * 1024**3 @@ -55,7 +57,7 @@ class SparseRefitTransport: """Transport-specific callbacks used by the shared streaming pipeline.""" - name: str + name: VllmRefitTransportName transfer_workers: int send: Callable[[bytes, int, int], dict[str, Any]] cleanup: Callable[[], None] @@ -165,13 +167,6 @@ def _request(self, method: str, key: str, body: bytes | None = None) -> Any: ) -def refit_env_int(name: str, *, default: int, min_value: int = 1) -> int: - value = int(os.getenv(name) or default) - if value < min_value: - raise ValueError(f"{name} must be >= {min_value}.") - return value - - def sparse_payload_checksum(body: bytes | bytearray) -> str: return hashlib.blake2b(body, digest_size=16).hexdigest() @@ -221,14 +216,11 @@ def _get_manifest_s3_store(bucket: str, region: str) -> _S3ObjectStore: def sparse_export_chunk_size( delta_tracker: DeltaCompressionTracker, - transport: str, + transport: VllmRefitTransportName, ) -> int: - default_mib = 64 if transport == "s3" else 256 - requested = refit_env_int( - f"NRL_REFIT_{transport.upper()}_EXPORT_CHUNK_BYTES", - default=default_mib * 1024**2, - min_value=1, - ) + requested = delta_tracker.refit_config.delta_compression.export_chunk_bytes[ + transport + ] if torch.cuda.is_available(): requested = min(requested, get_target_packed_tensor_size()) return min(requested, delta_tracker.sparse_bucket_size_bytes) @@ -245,7 +237,7 @@ def init_sparse_delta_baseline_from_iterator( delta_tracker: DeltaCompressionTracker, shard_rank: int, shard_count: int, - transport: str, + transport: VllmRefitTransportName, ) -> None: start_s = time.perf_counter() export_chunk_size = sparse_export_chunk_size(delta_tracker, transport) @@ -280,10 +272,8 @@ def stream_sparse_delta_payloads( shard_count: int, ) -> dict[str, int]: prefix = transport.name.upper() - encode_workers = refit_env_int( - f"NRL_REFIT_{prefix}_ENCODE_WORKERS", - default=max(2, min(8, os.cpu_count() or 8)), - ) + refit_config = delta_tracker.refit_config + encode_workers = refit_config.tuning.encode_workers[transport.name] encode_executor = _executor(f"refit-{transport.name}-encode", encode_workers) serialize_workers = min(4, encode_workers) serialize_executor = _executor( @@ -322,7 +312,10 @@ def serialize_payloads( raw_body = buffer.getvalue() serialize_s = time.perf_counter() - started started = time.perf_counter() - body = zstd_compress(raw_body, f"NRL_REFIT_{prefix}_ZSTD_THREADS") + body = zstd_compress( + raw_body, + refit_config.delta_compression.zstd_threads[transport.name], + ) compress_s = time.perf_counter() - started return ( body, @@ -524,14 +517,17 @@ def stream_sparse_delta_payloads_via_s3_manifest( endpoint_urls = vllm_refit_endpoints(refit_targets, G_VLLM_REFIT_S3_MANIFEST_PATH) if not endpoint_urls: raise ValueError("At least one vLLM S3 refit URL is required.") - bucket = os.getenv("NRL_REFIT_S3_BUCKET", "").strip() + refit_config = delta_tracker.refit_config + bucket = (refit_config.storage.s3_bucket or "").strip() if not bucket: - raise RuntimeError("NRL_REFIT_S3_BUCKET must be set for S3 refit.") - store = _get_manifest_s3_store( - bucket, - os.getenv("NRL_REFIT_S3_REGION", "us-east-1").strip() or "us-east-1", - ) - object_prefix = os.getenv("NRL_REFIT_S3_PREFIX", "nemo-rl-refit").strip("/") + raise RuntimeError( + "policy.generation.refit_cfg.storage.s3_bucket must be set for S3 refit." + ) + region = refit_config.storage.s3_region.strip() + if not region: + raise ValueError("refit_cfg.storage.s3_region must not be empty.") + store = _get_manifest_s3_store(bucket, region) + object_prefix = refit_config.storage.s3_prefix.strip("/") run_prefix = ( f"{object_prefix}/{transfer_id}/{shard_rank:06d}" if object_prefix @@ -590,10 +586,7 @@ def cleanup() -> None: delta_tracker=delta_tracker, transport=SparseRefitTransport( name="s3", - transfer_workers=refit_env_int( - "NRL_REFIT_S3_UPLOAD_WORKERS", - default=max(4, min(32, os.cpu_count() or 32)), - ), + transfer_workers=refit_config.tuning.transfer_workers["s3"], send=send, cleanup=cleanup, ), @@ -611,12 +604,12 @@ def download_s3_refit_payload( return decode_sparse_payload(body, str(manifest["checksum"])) -def zstd_compress(raw: bytes, threads_env: str) -> bytes: +def zstd_compress(raw: bytes, threads: int) -> bytes: compressor = getattr(_STREAM_LOCAL, "zstd_compressor", None) if compressor is None: compressor = zstandard.ZstdCompressor( level=1, - threads=refit_env_int(threads_env, default=0, min_value=0), + threads=threads, ) _STREAM_LOCAL.zstd_compressor = compressor return compressor.compress(raw) diff --git a/nemo_rl/utils/weight_transfer_zmq.py b/nemo_rl/utils/weight_transfer_zmq.py index 1ec44ebacef..c7f13577e28 100644 --- a/nemo_rl/utils/weight_transfer_zmq.py +++ b/nemo_rl/utils/weight_transfer_zmq.py @@ -14,6 +14,7 @@ """Transactional ZeroMQ value plane for remote sparse vLLM refit.""" +import hmac import json import threading import time @@ -26,6 +27,7 @@ import zmq +from nemo_rl.models.generation.vllm.config import VllmRefitTuningConfig from nemo_rl.utils.weight_transfer_http import ( merge_vllm_refit_metrics, vllm_refit_api_key, @@ -36,7 +38,6 @@ ) from nemo_rl.utils.weight_transfer_stream import ( SparseRefitTransport, - refit_env_int, sparse_payload_checksum, stream_sparse_delta_payloads, ) @@ -73,11 +74,13 @@ def __init__( *, timeout_s: float, producer_id: int, + retries: int, api_key: str | None = None, ) -> None: self._address = address self._timeout_ms = max(1, int(timeout_s * 1000)) self._producer_id = producer_id + self._retries = retries self._api_key = api_key self._socket = zmq.Context.instance().socket(zmq.DEALER) _configure_socket(self._socket, 2) @@ -113,16 +116,14 @@ def send_payload( if self._api_key is not None: metadata["api_key"] = self._api_key metadata_frame = _json_bytes(metadata) - retries = refit_env_int("NRL_REFIT_ZMQ_RETRIES", default=3, min_value=0) - - for attempt in range(retries + 1): + for attempt in range(self._retries + 1): try: self._socket.send_multipart( [_DATA, metadata_frame, body], copy=False, ) except zmq.Again: - if attempt == retries: + if attempt == self._retries: break continue @@ -151,7 +152,7 @@ def send_payload( return reply raise RuntimeError(f"ZeroMQ sparse refit rejected payload: {reply}") - if attempt < retries: + if attempt < self._retries: time.sleep(min(0.05 * 2**attempt, 0.5)) raise TimeoutError( @@ -172,19 +173,19 @@ def __init__( bind_address: str, api_key_env_var: str | None, timeout_s: float, + tuning: VllmRefitTuningConfig, ) -> None: self._apply_payload = apply_payload self._bind_address = bind_address self._token = vllm_refit_api_key(api_key_env_var) self._timeout_s = timeout_s + self._retries = tuning.zmq_retries self._stop = threading.Event() self._ready = threading.Event() self._thread: threading.Thread | None = None self._endpoint: str | None = None self._error: Exception | None = None - self._payload_workers = refit_env_int( - "NRL_REFIT_ZMQ_RELAY_PAYLOAD_WORKERS", default=16 - ) + self._payload_workers = tuning.zmq_relay_payload_workers self._transfer_lock = threading.Lock() self._transfer_condition = threading.Condition(self._transfer_lock) self._transfers: dict[str, _RelayTransfer] = {} @@ -197,7 +198,7 @@ def __init__( thread_name_prefix="nrl-zmq-payload", ) self._forward_executor = ThreadPoolExecutor( - max_workers=refit_env_int("NRL_REFIT_ZMQ_RELAY_FORWARD_WORKERS", default=8), + max_workers=tuning.zmq_relay_forward_workers, thread_name_prefix="nrl-zmq-forward", ) @@ -318,6 +319,7 @@ def _forward( address, timeout_s=self._timeout_s, producer_id=0, + retries=self._retries, api_key=self._token, ) clients[address] = client @@ -359,7 +361,11 @@ def _parse_data_message( metadata = json.loads(raw_metadata) if metadata.get("protocol") != _PROTOCOL: raise ValueError("Unsupported ZeroMQ sparse refit protocol.") - if self._token is not None and metadata.get("api_key") != self._token: + supplied_token = metadata.get("api_key") + if self._token is not None and ( + not isinstance(supplied_token, str) + or not hmac.compare_digest(self._token, supplied_token) + ): raise PermissionError("ZeroMQ sparse refit producer authentication failed.") transfer_id = str(metadata["transfer_id"]) producer_id = int(metadata["producer_id"]) @@ -494,6 +500,7 @@ def stream_sparse_delta_payloads_via_zmq( if not addresses: raise ValueError("At least one ZeroMQ sparse refit address is required.") address = addresses[shard_rank % len(addresses)] + tuning = delta_tracker.refit_config.tuning api_key = vllm_refit_api_key(api_key_env_var) local = threading.local() @@ -506,6 +513,7 @@ def send( address, timeout_s=timeout_s, producer_id=shard_rank, + retries=tuning.zmq_retries, api_key=api_key, ) local.client = client @@ -533,7 +541,7 @@ def cleanup() -> None: delta_tracker=delta_tracker, transport=SparseRefitTransport( name="zmq", - transfer_workers=refit_env_int("NRL_REFIT_ZMQ_SEND_WORKERS", default=4), + transfer_workers=tuning.transfer_workers["zmq"], send=send, cleanup=cleanup, ), diff --git a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py index 66936b06c38..eda7045d9cd 100644 --- a/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py +++ b/nemo_rl/weight_sync/vllm_remote_sparse_weight_synchronizer.py @@ -24,7 +24,7 @@ from nemo_rl.models.generation.vllm.config import ( VllmConfig, - VllmDeltaCompressionConfig, + normalize_vllm_refit_config, ) from nemo_rl.utils.timer import Timer from nemo_rl.utils.weight_transfer_http import ( @@ -57,23 +57,26 @@ def validate_vllm_remote_sparse_refit( if transport not in _REMOTE_SPARSE_TRANSPORTS: raise ValueError(f"Unsupported vLLM refit transport {transport!r}.") vllm_cfg = config["vllm_cfg"] - delta_config = config.get("delta_compression") + refit_config = normalize_vllm_refit_config(config) + assert refit_config is not None if ( colocated or not megatron_enabled or vllm_cfg["precision"] == "fp8" or vllm_cfg["kv_cache_dtype"].startswith("fp8") - or delta_config is None or config.get("quant_cfg") or config.get("real_quant") ): raise ValueError( f"{transport} requires a non-colocated Megatron policy, BF16/FP16 " - "vLLM, delta compression, and an unquantized rollout." + "vLLM, and an unquantized rollout." + ) + if transport == "vllm_s3_sparse" and not ( + refit_config.storage.s3_bucket and refit_config.storage.s3_bucket.strip() + ): + raise ValueError( + "vllm_s3_sparse requires policy.generation.refit_cfg.storage.s3_bucket." ) - config["delta_compression"] = VllmDeltaCompressionConfig.model_validate( - delta_config - ) return _REMOTE_SPARSE_TRANSPORTS[transport] @@ -99,6 +102,7 @@ def __init__( self._baseline_init_refs = list(baseline_init_refs or ()) self._baseline_commit_refs: list[Any] = [] self._stale = True + self._poisoned = False def sync_weights( self, @@ -106,6 +110,15 @@ def sync_weights( timer: Timer | None = None, kv_scales: dict[str, float] | None = None, ) -> dict[str, float]: + if self._poisoned: + raise RuntimeError( + "Sparse refit synchronizer was poisoned by a prior failed sync: " + "receivers may be partially applied while the trainer baseline is " + "uncommitted, so re-applying deltas would double-XOR and corrupt " + "weights. Reload the rollout workers from a known-good checkpoint " + "and re-initialize the communicator before syncing again. See " + "https://github.com/NVIDIA-NeMo/RL/issues/3274." + ) context = ( timer.time("prepare_for_generation/transfer_and_update_weights") if timer @@ -219,6 +232,7 @@ def sync_weights( succeeded = True finally: if not succeeded: + self._poisoned = True if self._transport == "zmq" and not relay_flushed: with suppress(Exception): self._request_receivers( diff --git a/tests/test_suites/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.sh b/tests/test_suites/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-deltaweight-noncolocated.sh similarity index 100% rename from tests/test_suites/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.sh rename to tests/test_suites/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-deltaweight-noncolocated.sh diff --git a/tests/test_suites/nightly.txt b/tests/test_suites/nightly.txt index 95db3bfbf34..a30f4f7c862 100644 --- a/tests/test_suites/nightly.txt +++ b/tests/test_suites/nightly.txt @@ -93,7 +93,7 @@ tests/test_suites/llm/grpo-qwen3-8b-base-dapo-2n8g-long-megatron-qa-nvfp4-w4a16. # Non-colocated tests/test_suites/llm/grpo-llama3.1-8b-instruct-2n8g-fsdp2tp1-noncolocated.sh tests/test_suites/llm/grpo-llama3.2-1b-instruct-2n8g-megatron_generation-noncolocated.sh -tests/test_suites/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-noncolocated.sh +tests/test_suites/llm/grpo-qwen3-30ba3b-4n8g-megatron-zmq-deltaweight-noncolocated.sh # Nemotron Super 49B #https://github.com/NVIDIA-NeMo/RL/issues/1374 diff --git a/tests/unit/models/generation/test_vllm_sparse_refit.py b/tests/unit/models/generation/test_vllm_sparse_refit.py index e0e82d1b645..1bee0dfa7ea 100644 --- a/tests/unit/models/generation/test_vllm_sparse_refit.py +++ b/tests/unit/models/generation/test_vllm_sparse_refit.py @@ -491,7 +491,7 @@ def test_sparse_refit_api_auth_dispatch_and_error_mapping(monkeypatch) -> None: def test_sync_sparse_refit_server_shutdown_cleans_transport_resources( - monkeypatch, + monkeypatch, caplog ) -> None: import uvicorn @@ -519,6 +519,7 @@ def run(self) -> None: monkeypatch.setattr(uvicorn, "Server", Server) monkeypatch.setattr(refit_module, "_get_free_port_local", lambda *_args: 12345) monkeypatch.setattr(refit_module, "_get_node_ip_local", lambda: "10.0.0.1") + refit_module._warn_unauthenticated_refit_server.cache_clear() config = { "vllm_cfg": {"async_engine": False, "http_refit_server_port": None}, "port_range_low": 10000, @@ -534,6 +535,7 @@ def run(self) -> None: assert configs[0].port == 12345 assert servers[0].ran.wait(timeout=1.0) assert receiver.report_refit_server_base_url() == "http://10.0.0.1:12345" + assert "binding 0.0.0.0 without an API key" in caplog.text relay = MagicMock() receiver._zmq_refit_server = (relay, "tcp://relay") @@ -593,6 +595,7 @@ def test_async_sparse_refit_exposes_zmq_relay(monkeypatch) -> None: bind_address="tcp://0.0.0.0:12345", api_key_env_var=None, timeout_s=600.0, + tuning=receiver._refit_config.tuning, ) server.start.assert_called_once_with() receiver.configure_zmq_sparse_refit_relay( diff --git a/tests/unit/models/policy/test_megatron_remote_sparse_refit.py b/tests/unit/models/policy/test_megatron_remote_sparse_refit.py index f3d5c62a7a4..8b1b697cb9b 100644 --- a/tests/unit/models/policy/test_megatron_remote_sparse_refit.py +++ b/tests/unit/models/policy/test_megatron_remote_sparse_refit.py @@ -16,13 +16,17 @@ import torch -from nemo_rl.models.generation.vllm.config import VllmDeltaCompressionConfig +from nemo_rl.models.generation.vllm.config import VllmRefitConfig from nemo_rl.models.policy.workers.megatron_remote_sparse_refit import ( MegatronRemoteSparseRefit, ) -_DELTA_CONFIG = VllmDeltaCompressionConfig( - encoding="overwrite", sparse_bucket_size_bytes=1024 +_REFIT_CONFIG = VllmRefitConfig( + delta_compression={ + "encoding": "overwrite", + "sparse_bucket_size_bytes": 1024, + }, + baseline={"in_memory": True}, ) @@ -33,13 +37,12 @@ def export(): return SimpleNamespace(_iter_params_with_optional_kv_scales=export) -def test_remote_sparse_initializes_canonical_hf_baseline(monkeypatch) -> None: - monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") +def test_remote_sparse_initializes_canonical_hf_baseline() -> None: weights = [ ("embedding.weight", torch.ones(2, 3)), ("linear.weight", torch.ones(4, 3)), ] - remote_refit = MegatronRemoteSparseRefit(_worker(weights), _DELTA_CONFIG) + remote_refit = MegatronRemoteSparseRefit(_worker(weights), _REFIT_CONFIG) info = remote_refit.initialize_baseline( shard_rank=0, shard_count=1, transport="zmq" @@ -53,7 +56,14 @@ def test_remote_sparse_initializes_canonical_hf_baseline(monkeypatch) -> None: def test_remote_sparse_preserves_xor_config() -> None: remote_refit = MegatronRemoteSparseRefit( - _worker(), _DELTA_CONFIG.model_copy(update={"encoding": "xor"}) + _worker(), + VllmRefitConfig( + delta_compression={ + "encoding": "xor", + "sparse_bucket_size_bytes": 1024, + }, + baseline={"in_memory": True}, + ), ) assert remote_refit._tracker.encoding == "xor" @@ -61,7 +71,7 @@ def test_remote_sparse_preserves_xor_config() -> None: def test_remote_sparse_streams_one_canonical_path_and_drains_cuda(monkeypatch) -> None: weights = [("model.weight", torch.ones(2))] - remote_refit = MegatronRemoteSparseRefit(_worker(weights), _DELTA_CONFIG) + remote_refit = MegatronRemoteSparseRefit(_worker(weights), _REFIT_CONFIG) expected = {"payloads": 1, "changed_elements": 2, "total_elements": 2} events = [] @@ -104,7 +114,7 @@ def stream(iterator, **kwargs): def test_remote_sparse_finishes_single_tracker(monkeypatch) -> None: - remote_refit = MegatronRemoteSparseRefit(_worker(), _DELTA_CONFIG) + remote_refit = MegatronRemoteSparseRefit(_worker(), _REFIT_CONFIG) events = [] monkeypatch.setattr( remote_refit._tracker, diff --git a/tests/unit/reference_configs/grpo_math_1B.yaml b/tests/unit/reference_configs/grpo_math_1B.yaml index 9dcabd2fc79..39aa2459723 100644 --- a/tests/unit/reference_configs/grpo_math_1B.yaml +++ b/tests/unit/reference_configs/grpo_math_1B.yaml @@ -333,7 +333,7 @@ policy: stop_token_ids: null stop_strings: null refit_transport: null # Set to "vllm_s3_sparse" or "vllm_zmq_sparse" for remote sparse-delta refit. - delta_compression: null # Set {} for XOR and the 512 MiB bucket; null uses the existing refit path. + refit_cfg: null # Optional tuning and storage settings for remote sparse-delta refit. mcore_generation_config: async_engine: false max_model_len: ${policy.max_total_sequence_length} diff --git a/tests/unit/utils/test_weight_transfer_stream.py b/tests/unit/utils/test_weight_transfer_stream.py index 9acc9885b51..4ab661fc657 100644 --- a/tests/unit/utils/test_weight_transfer_stream.py +++ b/tests/unit/utils/test_weight_transfer_stream.py @@ -17,12 +17,14 @@ import threading from concurrent.futures import ThreadPoolExecutor from types import SimpleNamespace +from typing import Any import pytest import requests import torch import zstandard +from nemo_rl.models.generation.vllm.config import VllmRefitConfig from nemo_rl.utils import ( weight_transfer_http, weight_transfer_stream, @@ -45,8 +47,59 @@ ) +def _refit_config( + *, + encoding: str = "overwrite", + bucket_bytes: int = 1024, + verify_samples: int = 0, + s3_bucket: str | None = None, + s3_region: str = "us-east-1", + s3_prefix: str = "nemo-rl-refit", + s3_export_bytes: int = 64 * 1024**2, + zmq_export_bytes: int = 256 * 1024**2, + s3_encode_workers: int = 8, + zmq_encode_workers: int = 8, + s3_transfer_workers: int = 32, + zmq_transfer_workers: int = 4, + zmq_retries: int = 3, +) -> VllmRefitConfig: + return VllmRefitConfig( + delta_compression={ + "encoding": encoding, + "sparse_bucket_size_bytes": bucket_bytes, + "export_chunk_bytes": { + "s3": s3_export_bytes, + "zmq": zmq_export_bytes, + }, + }, + storage={ + "s3_bucket": s3_bucket, + "s3_region": s3_region, + "s3_prefix": s3_prefix, + }, + baseline={"in_memory": True}, + tuning={ + "encode_workers": { + "s3": s3_encode_workers, + "zmq": zmq_encode_workers, + }, + "transfer_workers": { + "s3": s3_transfer_workers, + "zmq": zmq_transfer_workers, + }, + "zmq_retries": zmq_retries, + }, + verify_samples_per_payload=verify_samples, + ) + + class _SparsePipelineTracker: sparse_bucket_size_bytes = 1 + refit_config = _refit_config( + bucket_bytes=1, + zmq_export_bytes=1, + zmq_encode_workers=1, + ) @staticmethod def prepare_sparse_delta_payload(chunk): @@ -68,6 +121,7 @@ def prepare_sparse_delta_payload(chunk): class _BaselineNamesTracker: sparse_bucket_size_bytes = 4 + refit_config = _refit_config(bucket_bytes=4, zmq_export_bytes=4) def __init__(self) -> None: self.names = [] @@ -96,10 +150,10 @@ def _stream_sparse_test_payloads(tensors, send_payload): assert cleaned -def _delta_tracker(encoding: str = "overwrite") -> DeltaCompressionTracker: - return DeltaCompressionTracker( - {"encoding": encoding, "sparse_bucket_size_bytes": 1024} - ) +def _delta_tracker( + encoding: str = "overwrite", **config: Any +) -> DeltaCompressionTracker: + return DeltaCompressionTracker(_refit_config(encoding=encoding, **config)) def _baseline_names(tensors, *, rank: int): @@ -114,11 +168,6 @@ def _baseline_names(tensors, *, rank: int): return tracker.names -@pytest.fixture(autouse=True) -def _in_memory_baseline(monkeypatch) -> None: - monkeypatch.setenv("NRL_REFIT_BASELINE_IN_MEMORY", "1") - - def test_delta_tracker_commits_only_successful_syncs() -> None: tracker = _delta_tracker() tensor = torch.tensor([1.0, 2.0, 3.0]) @@ -132,9 +181,8 @@ def test_delta_tracker_commits_only_successful_syncs() -> None: assert not tracker.prepare_sparse_delta_payload([("weight", tensor)])[0][2] -def test_delta_tracker_emits_bounded_verification_budget(monkeypatch) -> None: - monkeypatch.setenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "2") - tracker = _delta_tracker() +def test_delta_tracker_emits_bounded_verification_budget() -> None: + tracker = _delta_tracker(verify_samples=2) tensor = torch.tensor([1.0, 2.0, 3.0, 4.0]) tracker.snapshot_baseline([("weight", tensor)]) tensor[[1, 3]] += 1 @@ -221,12 +269,29 @@ def test_sparse_index_encoding_preserves_uint64_locations() -> None: assert torch.equal(decoded, locations) +def test_sparse_index_encoding_preserves_uint32_locations() -> None: + locations = torch.tensor([0, 2**16 + 1]) + packed, _, metadata = encode_sparse_infos( + [ + ( + "weight", + torch.empty(2), + locations, + torch.ones(2, dtype=torch.int32), + "overwrite", + ) + ], + ) + + assert metadata[0]["index_encoding"] == "deltas" + assert packed.numel() == 2 * 4 + decoded = sparse_locations_for_item(metadata[0], packed, device="cpu") + assert torch.equal(decoded, locations) + + @pytest.mark.parametrize("encoding", ["xor", "overwrite"]) -def test_delta_tracker_encodes_fp8_weight_and_scale_bits( - monkeypatch, encoding: str -) -> None: - monkeypatch.setenv("NRL_REFIT_VERIFY_SAMPLES_PER_PAYLOAD", "2") - tracker = _delta_tracker(encoding) +def test_delta_tracker_encodes_fp8_weight_and_scale_bits(encoding: str) -> None: + tracker = _delta_tracker(encoding, verify_samples=2) weight = torch.tensor([0x38, 0x40, 0x48], dtype=torch.uint8).view( torch.float8_e4m3fn ) @@ -260,8 +325,8 @@ def test_delta_tracker_encodes_fp8_weight_and_scale_bits( )[0][2] -def test_delta_tracker_rejects_arithmetic_encoding() -> None: - with pytest.raises(ValueError, match="Unsupported sparse-refit operation"): +def test_refit_config_rejects_arithmetic_encoding() -> None: + with pytest.raises(ValueError, match="Input should be 'xor' or 'overwrite'"): _delta_tracker("add") @@ -318,9 +383,7 @@ def test_refit_http_error_preserves_non_json_status_and_body(monkeypatch) -> Non ) -def test_sparse_export_finishes_before_blocked_transfers(monkeypatch) -> None: - monkeypatch.setenv("NRL_REFIT_ZMQ_ENCODE_WORKERS", "1") - monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "1") +def test_sparse_export_finishes_before_blocked_transfers() -> None: exported = threading.Event() release_transfers = threading.Event() result = [] @@ -349,9 +412,7 @@ def run(): assert result == [{"payloads": 4, "changed_elements": 4, "total_elements": 4}] -def test_sparse_export_finishes_before_transfer_error(monkeypatch) -> None: - monkeypatch.setenv("NRL_REFIT_ZMQ_ENCODE_WORKERS", "1") - monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "1") +def test_sparse_export_finishes_before_transfer_error() -> None: exported = [] def tensors(): @@ -368,10 +429,7 @@ def fail_transfer(_body, _payload_index): assert exported == list(range(4)) -def test_sparse_transport_cleanup_runs_on_transfer_workers(monkeypatch) -> None: - monkeypatch.setenv("NRL_REFIT_ZMQ_ENCODE_WORKERS", "1") - monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "1") - +def test_sparse_transport_cleanup_runs_on_transfer_workers() -> None: class Transport: name = "zmq" transfer_workers = 2 @@ -402,11 +460,12 @@ def cleanup(self) -> None: assert transport.send_threads == transport.cleanup_threads -def test_sparse_stream_coalesces_export_chunks(monkeypatch) -> None: - monkeypatch.setenv("NRL_REFIT_ZMQ_ENCODE_WORKERS", "2") - monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "4") - tracker = _delta_tracker() - tracker.sparse_bucket_size_bytes = 8 +def test_sparse_stream_coalesces_export_chunks() -> None: + tracker = _delta_tracker( + bucket_bytes=8, + zmq_export_bytes=4, + zmq_encode_workers=2, + ) tensors = [(f"weight-{index}", torch.zeros(1)) for index in range(4)] tracker.snapshot_baseline(tensors) for _, tensor in tensors: @@ -454,28 +513,21 @@ def send(body, payload_index): ].tolist() == [1065353216] -def test_sparse_export_chunk_defaults_are_transport_specific(monkeypatch) -> None: - monkeypatch.delenv("NRL_REFIT_S3_EXPORT_CHUNK_BYTES", raising=False) - monkeypatch.delenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", raising=False) - tracker = _delta_tracker() - tracker.sparse_bucket_size_bytes = 1024**3 +def test_sparse_export_chunk_defaults_are_transport_specific() -> None: + tracker = _delta_tracker(bucket_bytes=1024**3) assert sparse_export_chunk_size(tracker, "s3") == 64 * 1024**2 assert sparse_export_chunk_size(tracker, "zmq") == 256 * 1024**2 -def test_sparse_baseline_snapshots_only_owned_export_chunks( - monkeypatch, capsys -) -> None: - monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "4") +def test_sparse_baseline_snapshots_only_owned_export_chunks(capsys) -> None: tensors = [(f"weight-{index}", torch.tensor([float(index)])) for index in range(4)] assert _baseline_names(tensors, rank=1) == ["weight-1", "weight-3"] assert "chunks=4" in capsys.readouterr().out -def test_sparse_stream_sends_only_owned_export_chunks(monkeypatch) -> None: - monkeypatch.setenv("NRL_REFIT_ZMQ_EXPORT_CHUNK_BYTES", "1") +def test_sparse_stream_sends_only_owned_export_chunks() -> None: sent = [] transport = SparseRefitTransport( "zmq", @@ -499,23 +551,23 @@ def test_sparse_stream_sends_only_owned_export_chunks(monkeypatch) -> None: def test_s3_manifest_transport_validates_configuration(monkeypatch) -> None: + tracker = SimpleNamespace(refit_config=_refit_config(s3_bucket="bucket")) kwargs = { "iterator": (), - "delta_tracker": SimpleNamespace(), + "delta_tracker": tracker, "transfer_id": "transfer", "api_key_env_var": None, "timeout_s": 1.0, "shard_rank": 0, "shard_count": 1, } - monkeypatch.setenv("NRL_REFIT_S3_BUCKET", "bucket") with pytest.raises(ValueError, match="URL is required"): weight_transfer_stream.stream_sparse_delta_payloads_via_s3_manifest( refit_targets=[], **kwargs ) - monkeypatch.delenv("NRL_REFIT_S3_BUCKET") - with pytest.raises(RuntimeError, match="NRL_REFIT_S3_BUCKET"): + tracker.refit_config = _refit_config() + with pytest.raises(RuntimeError, match="refit_cfg.storage.s3_bucket"): weight_transfer_stream.stream_sparse_delta_payloads_via_s3_manifest( refit_targets=["http://receiver"], **kwargs ) @@ -536,10 +588,6 @@ def delete(self, key) -> None: operations.append(("delete", key)) store = Store() - monkeypatch.setenv("NRL_REFIT_S3_BUCKET", store.bucket) - monkeypatch.setenv("NRL_REFIT_S3_REGION", store.region) - monkeypatch.setenv("NRL_REFIT_S3_PREFIX", "/prefix/") - monkeypatch.setenv("NRL_REFIT_S3_UPLOAD_WORKERS", "3") monkeypatch.setenv("NRL_TEST_REFIT_KEY", "secret") monkeypatch.setattr( weight_transfer_stream, @@ -569,7 +617,14 @@ def stream(iterator, **kwargs): result = weight_transfer_stream.stream_sparse_delta_payloads_via_s3_manifest( [("weight", torch.tensor([1.0]))], - delta_tracker=SimpleNamespace(), + delta_tracker=SimpleNamespace( + refit_config=_refit_config( + s3_bucket=store.bucket, + s3_region=store.region, + s3_prefix="/prefix/", + s3_transfer_workers=3, + ) + ), refit_targets=[" http://receiver-a/ ", "http://receiver-b"], transfer_id="transfer", api_key_env_var="NRL_TEST_REFIT_KEY", @@ -604,7 +659,6 @@ def test_zmq_stream_routes_shards_and_closes_clients(monkeypatch) -> None: sent = [] closed = [] monkeypatch.setenv("NRL_TEST_REFIT_KEY", "secret") - monkeypatch.setenv("NRL_REFIT_ZMQ_SEND_WORKERS", "2") class Client: def __init__(self, address, **kwargs) -> None: @@ -646,7 +700,9 @@ def stream(_iterator, **kwargs): kwargs = { "iterator": (), - "delta_tracker": SimpleNamespace(), + "delta_tracker": SimpleNamespace( + refit_config=_refit_config(zmq_transfer_workers=2) + ), "refit_targets": ["tcp://receiver-a", " tcp://receiver-b "], "transfer_id": "transfer", "api_key_env_var": "NRL_TEST_REFIT_KEY", @@ -666,7 +722,12 @@ def stream(_iterator, **kwargs): assert created == 2 * [ ( "tcp://receiver-b", - {"timeout_s": 7.0, "producer_id": 3, "api_key": "secret"}, + { + "timeout_s": 7.0, + "producer_id": 3, + "retries": 3, + "api_key": "secret", + }, ) ] assert closed == [True, True] @@ -701,6 +762,7 @@ def client(socket) -> ZmqSparseRefitClient: result._address = "tcp://receiver" result._timeout_ms = 1000 result._producer_id = 4 + result._retries = 1 result._api_key = "secret" result._socket = socket return result @@ -728,7 +790,6 @@ def client(socket) -> ZmqSparseRefitClient: with pytest.raises(RuntimeError, match="denied"): _send_zmq_payload(client(denied), 7, b"body") - monkeypatch.setenv("NRL_REFIT_ZMQ_RETRIES", "1") with pytest.raises(TimeoutError, match="payload 7"): _send_zmq_payload(client(Socket(send_failures=2)), 7, b"body") @@ -800,6 +861,7 @@ def apply_payload(body, metadata): bind_address="tcp://127.0.0.1:*", api_key_env_var="NRL_TEST_REFIT_KEY", timeout_s=5.0, + tuning=_refit_config().tuning, ) for items in received ] @@ -810,11 +872,13 @@ def apply_payload(body, metadata): addresses[0], timeout_s=5.0, producer_id=2, + retries=3, ) client = ZmqSparseRefitClient( addresses[0], timeout_s=5.0, producer_id=3, + retries=3, api_key="secret", ) body = b"compressed sparse payload" 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 c7078ac91fd..0c6f5a6c461 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 @@ -16,7 +16,7 @@ import pytest -from nemo_rl.models.generation.vllm.config import VllmDeltaCompressionConfig +from nemo_rl.models.generation.vllm.config import VllmRefitConfig from nemo_rl.weight_sync.vllm_remote_sparse_weight_synchronizer import ( VllmRemoteSparseWeightSynchronizer, validate_vllm_remote_sparse_refit, @@ -70,7 +70,10 @@ def _remote_sparse_sync( def _valid_config() -> dict: return { "refit_transport": "vllm_s3_sparse", - "delta_compression": {"encoding": "overwrite"}, + "refit_cfg": { + "delta_compression": {"encoding": "overwrite"}, + "storage": {"s3_bucket": "bucket"}, + }, "vllm_cfg": {"precision": "bfloat16", "kv_cache_dtype": "auto"}, } @@ -83,8 +86,9 @@ def test_validate_remote_sparse_refit_accepts_supported_scope(): ) == "s3" ) - assert config["delta_compression"] == VllmDeltaCompressionConfig( - encoding="overwrite" + assert config["refit_cfg"] == VllmRefitConfig( + delta_compression={"encoding": "overwrite"}, + storage={"s3_bucket": "bucket"}, ) @@ -94,7 +98,7 @@ def test_validate_remote_sparse_refit_accepts_supported_scope(): ({"refit_transport": "unknown"}, {}), ({}, {"colocated": True}), ({}, {"megatron_enabled": False}), - ({"delta_compression": None}, {}), + ({"refit_cfg": {"storage": {"s3_bucket": None}}}, {}), ({"quant_cfg": "fp8"}, {}), ({"vllm_cfg": {"precision": "fp8", "kv_cache_dtype": "auto"}}, {}), ( @@ -324,6 +328,9 @@ def test_failure_drains_receivers_without_committing_baseline(self, mock_ray, po with pytest.raises(RuntimeError, match="stream failed"): sync.sync_weights() + with pytest.raises(RuntimeError, match="poisoned by a prior failed sync"): + sync.sync_weights() + post.assert_called_once_with( ["http://receiver/nemo-rl/refit/flush"], {}, From cf6f49019edd51c2fcdca3e476e3c15b0a1bd006 Mon Sep 17 00:00:00 2001 From: Hollow Man Date: Fri, 17 Jul 2026 17:02:16 -0700 Subject: [PATCH 18/18] More detailed diagram in docs Signed-off-by: Hollow Man --- docs/design-docs/sparse-delta-refit.md | 222 +++++++++++++++++++------ 1 file changed, 175 insertions(+), 47 deletions(-) diff --git a/docs/design-docs/sparse-delta-refit.md b/docs/design-docs/sparse-delta-refit.md index 82f823bab74..830d7f1bcec 100644 --- a/docs/design-docs/sparse-delta-refit.md +++ b/docs/design-docs/sparse-delta-refit.md @@ -35,48 +35,91 @@ reject `refit_transport` during setup; extending them is tracked in ## Architecture ```mermaid -flowchart LR - SYNC["GRPO weight synchronizer"] - - subgraph P["Megatron policy cluster"] - MB["All policy ranks
Megatron Bridge HF export"] - OWN["Deterministic chunk owner
chunk_index modulo producers"] - BASE["Distributed canonical
CPU or mmap baseline"] - PIPE["Shared bounded pipeline
compare, encode, zstd"] - MB -->|canonical HF chunks| OWN - OWN -->|owned chunks| BASE - BASE -->|changed locations and bits| PIPE - end - - subgraph T["Value plane"] - S3["S3 objects
one upload per owned payload"] - Z0["ZeroMQ root relay
one cross-cluster send"] - ZT["Binary relay tree"] - Z0 --> ZT - end - - subgraph G["vLLM generation cluster"] - HTTP["HTTP control plane
prepare and global flush"] - RX["One receiver per inference node"] - STAGE["Compact payload staging
and bounded FIFO queue"] - RPC["Node-local collective RPC"] - APPLY["Each vLLM rank
reusable scratch + native load_weights()"] - HTTP -.-> RX - RX --> STAGE --> RPC --> APPLY +%%{init: {"flowchart": {"curve": "linear", "nodeSpacing": 24, "rankSpacing": 30, "diagramPadding": 8}}}%% +flowchart TB + subgraph SYSTEM[" "] + direction TB + SYNC["VllmRemoteSparseWeightSynchronizer
driver: orchestrates; no payload bytes"] + + subgraph ROW[" "] + direction LR + subgraph TRAIN["Megatron policy cluster"] + direction TB + subgraph PACTORS["MegatronPolicyWorker x N - Ray actors"] + direction TB + REMOTE["MegatronRemoteSparseRefit"] + EXPORT["Megatron Bridge HF export
all policy ranks participate
deterministic owner: chunk_index modulo producers"] + TRACKER["DeltaCompressionTracker + bounded pipeline
distributed CPU or mmap baseline
compare -> XOR/overwrite -> encode -> zstd"] + REMOTE --> EXPORT --> TRACKER + end + end + + subgraph VALUE["Cross-cluster value plane"] + direction TB + S3["S3 object
compressed payload bytes
one upload per owned payload"] + ZTREE["ZeroMQ binary relay tree
one cross-cluster root send per payload"] + end + + subgraph GEN["vLLM generation cluster"] + direction TB + subgraph GACTORS["VllmGenerationWorker x M - Ray actors; one receiver per node"] + direction TB + RECEIVER["VllmSparseRefitReceiver"] + API["FastAPI control plane
/prepare /s3-manifest /zmq-flush /flush"] + ZMQ["ZmqSparseRefitServer
relay node + staged-ACK tree"] + QUEUE["Deduplicating bounded FIFO
batch staging + single apply stream"] + RECEIVER --> API -->|downloaded S3 body| QUEUE + RECEIVER --> ZMQ -->|in-process callback| QUEUE + end + + subgraph VPROC["vLLM worker process x TP/PP/EP - reached by collective_rpc"] + APPLY["VllmInternalWorkerExtension -> VllmSparseDeltaApplier
decode -> reusable GPU scratch -> native load_weights()"] + end + QUEUE -->|path when ranks share a node; bytes otherwise| APPLY + end + + TRAIN ~~~ VALUE ~~~ GEN + end + + SYNC -. "Ray RPC: stream and finish(success)" .-> REMOTE + SYNC -. "HTTP: prepare, relay flush, global flush" .-> API + TRACKER -->|PUT compressed bytes| S3 + TRACKER -. "POST manifest pointer" .-> API + S3 -->|GET compressed bytes per receiver node| API + TRACKER -->|DEALER to ROUTER: DATA bytes| ZTREE + ZTREE -->|relay DATA bytes| ZMQ end - SYNC -. "prepare, flush, commit" .-> HTTP - SYNC -. "start stream" .-> PIPE - PIPE -->|compressed body| S3 - S3 -->|HTTP manifest, then GET| RX - PIPE -->|multipart DATA frame| Z0 - ZT -->|compressed body| RX - APPLY -. "success metrics" .-> SYNC - SYNC -. "commit exact source bits" .-> BASE + classDef coordinator fill:#f1edff,stroke:#7048e8,color:#3f248f,stroke-width:2px + classDef policy fill:#f2f8f2,stroke:#2f7d32,color:#245b27,stroke-width:1.5px + classDef s3 fill:#fff8e8,stroke:#9a6500,color:#704900,stroke-width:1.5px + classDef zmq fill:#edf9fa,stroke:#267783,color:#1e5962,stroke-width:1.5px + classDef receiver fill:#eef4ff,stroke:#2864dc,color:#17418f,stroke-width:1.5px + classDef queue fill:#fff1ef,stroke:#c43b31,color:#8d2721,stroke-width:1.5px + classDef apply fill:#f2f8f2,stroke:#2f7d32,color:#245b27,stroke-width:1.5px + class SYNC coordinator + class REMOTE,EXPORT,TRACKER policy + class S3 s3 + class ZTREE,ZMQ zmq + class RECEIVER,API receiver + class QUEUE queue + class APPLY apply + style SYSTEM fill:transparent,stroke:transparent + style ROW fill:transparent,stroke:transparent + style TRAIN fill:#f8fbf8,stroke:#2f7d32,stroke-dasharray:5 5 + style GEN fill:#f6f8ff,stroke:#2864dc,stroke-dasharray:5 5 ``` -*Figure 1. Control traffic is dashed. Payload data is solid. S3 and ZeroMQ -share every component except the value plane.* +*Figure 1. Dashed links are control messages or S3 pointers; solid links carry +payload bytes or invoke the node-local apply path. Amber components are S3, +teal components are ZeroMQ, and both transports share the sender pipeline, +receiver queue, native-loader apply path, verification, and baseline commit.* + +| Path | Cross-cluster transfer | Receiver handoff | +|---|---|---| +| S3 | One object upload per owned payload; each receiver gets a small manifest pointer and downloads the object | FastAPI callback -> deduplicating FIFO -> `collective_rpc` | +| ZeroMQ | One producer-to-root DATA send followed by binary-tree relay fanout | In-process relay callback -> the same FIFO -> `collective_rpc` | +| Within one node | No S3 or ZeroMQ hop | Staged path when ranks share storage; serialized bytes otherwise; no CUDA IPC | | Responsibility | Implementation | |---|---| @@ -89,6 +132,69 @@ share every component except the value plane.* | Queue receiver work and expose endpoints | [`vllm_sparse_refit.py`](../../nemo_rl/models/generation/vllm/vllm_sparse_refit.py) | | Apply canonical updates through native loaders | [`vllm_sparse_delta.py`](../../nemo_rl/models/generation/vllm/vllm_sparse_delta.py) | +The baseline advances only after the receiver version is globally committed: + +```mermaid +%%{init: {"flowchart": {"curve": "linear", "nodeSpacing": 24, "rankSpacing": 30, "diagramPadding": 8}}}%% +flowchart LR + subgraph VERSION0["Initialization"] + direction TB + INIT["Same checkpoint version W0
loaded independently on both clusters"] + T0["Trainer GPU
W0"] + B0["CPU or mmap baseline
B0 = W0"] + V0["vLLM GPU
W0; no baseline copy"] + INIT -. "local load" .-> T0 + INIT -. "local load" .-> V0 + T0 -->|canonical export to local baseline| B0 + end + + subgraph VERSION1["Refit after optimizer step 1"] + direction TB + T1["Trainer GPU
W1"] + D1["delta1
changed bits of W1 vs B0
XOR; overwrite where required"] + V1["vLLM GPU
apply delta1 -> W1"] + A1["All receiver ACKs
relay flush + apply flush + verify"] + B1["Baseline commit
B1 = exact W1 source bits"] + T1 --> D1 --> V1 --> A1 --> B1 + end + + subgraph VERSION2["Refit after optimizer step 2"] + direction TB + T2["Trainer GPU
W2"] + D2["delta2
changed bits of W2 vs B1
XOR; overwrite where required"] + V2["vLLM GPU
apply delta2 -> W2"] + A2["All receiver ACKs
relay flush + apply flush + verify"] + B2["Baseline commit
B2 = exact W2 source bits"] + T2 --> D2 --> V2 --> A2 --> B2 + end + + NEXT["..."] + + T0 -->|optimizer step 1| T1 -->|optimizer step 2| T2 --> NEXT + B0 -->|comparison reference| D1 + B1 -->|comparison reference| D2 + V0 -->|serve rollouts until commit| V1 -->|serve rollouts until commit| V2 --> NEXT + B0 -. "advance only after A1" .-> B1 + B1 -. "advance only after A2" .-> B2 + + classDef init fill:#f1edff,stroke:#7048e8,color:#3f248f,stroke-width:1.5px + classDef weight fill:#f2f8f2,stroke:#2f7d32,color:#245b27,stroke-width:1.5px + classDef delta fill:#fff8e8,stroke:#9a6500,color:#704900,stroke-width:1.5px + classDef ack fill:#eef4ff,stroke:#2864dc,color:#17418f,stroke-width:1.5px + class INIT init + class T0,T1,T2,B0,B1,B2,V0,V1,V2 weight + class D1,D2 delta + class A1,A2 ack + style VERSION0 fill:#f7f8fa,stroke:#c8ced8 + style VERSION1 fill:#f7f8fa,stroke:#c8ced8 + style VERSION2 fill:#f7f8fa,stroke:#c8ced8 +``` + +*Figure 2. One refit follows each configured training cadence, but only changed +canonical bytes cross the refit value plane. No full checkpoint is transferred +by this protocol at initialization or periodically. Every policy rank still +participates in the intra-cluster Megatron Bridge export.* + ## Protocol ### Baseline and ownership @@ -270,28 +376,50 @@ Each serialized payload is: ``` ```mermaid -flowchart TB +%%{init: {"flowchart": {"curve": "linear", "nodeSpacing": 24, "rankSpacing": 30, "diagramPadding": 8}}}%% +flowchart LR + CODEC["Shared encoder output
(locations, value groups, tensor metadata)"] + SERIAL["torch serialization + zstd
checksum over compressed body"] + CODEC --> SERIAL + subgraph S3F["S3 transport"] - SO["Object body
zstd-compressed torch payload"] - SM["HTTP manifest
bucket, region, key, checksum,
verification_candidates"] + direction TB + SO["S3 object body
compressed payload bytes"] + SM["HTTP manifest pointer
bucket, region, key, checksum,
transfer/producer/payload IDs, samples"] + SM -. "locates object" .-> SO end subgraph ZF["ZeroMQ DATA multipart"] + direction TB ZK["Frame 1: DATA"] - ZM["Frame 2: JSON metadata
transfer, producer, payload IDs,
checksum, samples, optional API key"] + ZM["Frame 2: JSON metadata
transfer/producer/payload IDs,
checksum, samples, optional API key"] ZB["Frame 3: zstd-compressed torch payload"] + ZK --> ZM --> ZB end - SO --> DECOMP["Checksum then zstd decode"] - SM -. "locates and authenticates object" .-> SO - ZK --> ZM --> ZB --> DECOMP - DECOMP --> PT["torch payload tuple"] + SERIAL -->|PUT bytes once| SO + SERIAL -->|DATA body| ZB + SERIAL -. "POST pointer" .-> SM + SO --> COMMON["Common receiver
checksum -> zstd -> torch deserialize"] + SM -. "identity + verification budget" .-> COMMON + ZB --> COMMON + ZM -. "identity + authentication" .-> COMMON + COMMON --> PT["Decoded payload tuple"] PT --> PI["packed location bytes
range or uint16/32/64 deltas"] PT --> PV["value groups by dtype
XOR or overwrite bits"] PT --> PM["per-tensor metadata
HF name, shape, offsets, operation"] + + classDef shared fill:#f1edff,stroke:#7048e8,color:#3f248f,stroke-width:1.5px + classDef s3 fill:#fff8e8,stroke:#9a6500,color:#704900,stroke-width:1.5px + classDef zmq fill:#edf9fa,stroke:#267783,color:#1e5962,stroke-width:1.5px + classDef decoded fill:#f2f8f2,stroke:#2f7d32,color:#245b27,stroke-width:1.5px + class CODEC,SERIAL shared + class SO,SM s3 + class ZK,ZM,ZB zmq + class COMMON,PT,PI,PV,PM decoded ``` -*Figure 2. S3 separates the object body from its control-plane manifest; +*Figure 3. S3 separates the object body from its control-plane manifest; ZeroMQ carries equivalent identity and body data as multipart frames. Both decode into the same codec tuple.*