diff --git a/requirements.txt b/requirements.txt index 8b016fb03..7e400e7ee 100644 --- a/requirements.txt +++ b/requirements.txt @@ -13,7 +13,7 @@ qwen_vl_utils # for VLM ray[default] ring_flash_attn sglang-router>=0.2.3 -vllm-router>=0.1.14 tensorboard transformers +vllm-router>=0.1.14 wandb diff --git a/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py b/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py index 56f53b6e5..150ea1573 100644 --- a/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py +++ b/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py @@ -4,7 +4,6 @@ import os import socket import time -import traceback from argparse import Namespace from collections.abc import Callable, Mapping, Sequence from typing import Any @@ -12,11 +11,11 @@ import ray import torch import torch.distributed as dist -import torch.multiprocessing as mp from megatron.core import mpu from ray import ObjectRef from ray.actor import ActorHandle from tqdm import tqdm +from vllm.distributed.weight_transfer.nccl_engine import NCCLTrainerSendWeightsArgs, NCCLWeightTransferEngine from slime.utils.distributed_utils import get_gloo_group @@ -26,179 +25,18 @@ logger = logging.getLogger(__name__) -# --------------------------------------------------------------------------- -# NcclBridge: isolate vLLM's PyNcclCommunicator in a subprocess so that it -# never coexists with torch.distributed NCCL groups in the Megatron trainer. -# -# vLLM's weight transfer uses raw NCCL (PyNcclCommunicator) which conflicts -# with torch.distributed's NCCL backend when both exist in the same process -# (see https://github.com/vllm-project/vllm/issues/5477). SGLang avoids -# this because it uses torch.distributed process groups for weight sync. -# --------------------------------------------------------------------------- - - -def _nccl_bridge_worker( - conn, - master_address: str, - master_port: int, - world_size: int, - device: int, - cvd: str, - env_snapshot: dict[str, str], -) -> None: - """Subprocess entry-point: creates PyNcclCommunicator and serves requests. - - GPU tensors are shared from the parent via CUDA IPC (torch.multiprocessing - handles this transparently). No GPU→CPU→GPU copies are needed. - - Protocol over *conn* (multiprocessing.Connection): - parent → child: - {"op": "broadcast", "tensors": [gpu_tensor, ...]} - {"op": "send_packed", "named_tensors": [(name, gpu_tensor), ...]} - None → shutdown - child → parent: - "ready" (after init) - "ok" (after each op) - "error: ..." - """ - try: - os.environ.update(env_snapshot) - if cvd: - os.environ["CUDA_VISIBLE_DEVICES"] = cvd - - import torch as _torch # noqa: PLC0415 — subprocess needs fresh import - import torch.multiprocessing # noqa: F401, PLC0415 — register CUDA IPC reducers - - _torch.cuda.set_device(device) - - from vllm.distributed.device_communicators.pynccl import PyNcclCommunicator # noqa: PLC0415 - from vllm.distributed.utils import StatelessProcessGroup # noqa: PLC0415 - - pg = StatelessProcessGroup.create( - host=master_address, - port=master_port, - rank=0, - world_size=world_size, - ) - comm = PyNcclCommunicator(pg, device=device) - - conn.send("ready") - - while True: - cmd = conn.recv() - if cmd is None: - break - - op = cmd["op"] - if op == "broadcast": - for t in cmd["tensors"]: - comm.broadcast(t, src=0, stream=_torch.cuda.current_stream()) - _torch.cuda.synchronize() - conn.send("ok") - - elif op == "send_packed": - # Prefer NCCLWeightTransferEngine.trainer_send_weights (newer vLLM). Some pip builds omit - # NCCLTrainerSendWeightsArgs but still ship packed_broadcast_producer. - try: - from vllm.distributed.weight_transfer.nccl_engine import ( # noqa: PLC0415 - NCCLTrainerSendWeightsArgs, - NCCLWeightTransferEngine, - ) - - trainer_args = NCCLTrainerSendWeightsArgs( - group=comm, - packed=True, - ) - NCCLWeightTransferEngine.trainer_send_weights( - iterator=iter(cmd["named_tensors"]), - trainer_args=trainer_args, - ) - except ImportError: - from vllm.distributed.weight_transfer.packed_tensor import ( # noqa: PLC0415 - DEFAULT_PACKED_BUFFER_SIZE_BYTES, - DEFAULT_PACKED_NUM_BUFFERS, - packed_broadcast_producer, - ) - - packed_broadcast_producer( - iterator=iter(cmd["named_tensors"]), - group=comm, - src=0, - post_iter_func=lambda x: x[1], - buffer_size_bytes=DEFAULT_PACKED_BUFFER_SIZE_BYTES, - num_buffers=DEFAULT_PACKED_NUM_BUFFERS, - ) - _torch.cuda.synchronize() - conn.send("ok") - - except Exception as e: - try: - conn.send(f"error: {e}") - except Exception: - pass - traceback.print_exc() +def _begin_vllm_weight_update_session(rollout_engines: Sequence[ActorHandle]) -> None: + if dist.get_rank() == 0: + logger.info("vLLM weight update: start_weight_update") + ray.get([engine.start_weight_update.remote(is_checkpoint_format=True) for engine in rollout_engines]) + dist.barrier(group=get_gloo_group()) -class _NcclBridge: - """Runs vLLM's PyNcclCommunicator in a separate subprocess. - - This prevents NCCL communicator conflicts with torch.distributed groups - that already exist in the Megatron trainer process. GPU tensors are shared - with the subprocess via CUDA IPC (handled transparently by - torch.multiprocessing), avoiding any GPU→CPU→GPU copies. - """ - - def __init__(self, master_address: str, master_port: int, world_size: int, device: int): - ctx = mp.get_context("spawn") - self._parent_conn, child_conn = ctx.Pipe() - - env_snapshot = dict(os.environ) - cvd = os.environ.get("CUDA_VISIBLE_DEVICES", "") - - self._process = ctx.Process( - target=_nccl_bridge_worker, - args=(child_conn, master_address, master_port, world_size, device, cvd, env_snapshot), - daemon=True, - ) - self._process.start() - - msg = self._parent_conn.recv() - if isinstance(msg, str) and msg.startswith("error:"): - raise RuntimeError(f"NcclBridge init failed: {msg}") - if msg != "ready": - raise RuntimeError(f"NcclBridge init unexpected response: {msg}") - logger.info("NcclBridge ready (pid=%d, device=%d)", self._process.pid, device) - - def broadcast_tensors(self, tensors: list[torch.Tensor]) -> None: - """Broadcast a list of tensors (one-by-one) via the bridge subprocess.""" - gpu_tensors = [t.contiguous() for t in tensors] - self._parent_conn.send({"op": "broadcast", "tensors": gpu_tensors}) - self._wait_ok("broadcast_tensors") - - def send_weights_packed(self, named_tensors: list[tuple[str, torch.Tensor]]) -> None: - """Send weights using vLLM's packed broadcast protocol.""" - gpu_pairs = [] - for name, t in named_tensors: - data = t.data if hasattr(t, "data") else t - gpu_pairs.append((name, data.contiguous())) - self._parent_conn.send({"op": "send_packed", "named_tensors": gpu_pairs}) - self._wait_ok("send_weights_packed") - - def _wait_ok(self, label: str, timeout: float = 600.0) -> None: - if not self._parent_conn.poll(timeout): - raise TimeoutError(f"NcclBridge {label} timed out after {timeout}s") - msg = self._parent_conn.recv() - if msg != "ok": - raise RuntimeError(f"NcclBridge {label} failed: {msg}") - - def shutdown(self) -> None: - try: - self._parent_conn.send(None) - self._process.join(timeout=30) - except Exception: - pass - if self._process.is_alive(): - self._process.terminate() +def _end_vllm_weight_update_session(rollout_engines: Sequence[ActorHandle]) -> None: + if dist.get_rank() == 0: + logger.info("vLLM weight update: finish_weight_update") + ray.get([engine.finish_weight_update.remote() for engine in rollout_engines]) + dist.barrier(group=get_gloo_group()) class UpdateWeightFromDistributed: @@ -290,6 +128,25 @@ def update_weights(self) -> None: ) dist.barrier(group=get_gloo_group()) + _begin_vllm_weight_update_session(self.rollout_engines) + try: + self._sync_weights_to_rollout_engines() + finally: + _end_vllm_weight_update_session(self.rollout_engines) + + dist.barrier(group=get_gloo_group()) + if dist.get_rank() == 0: + # int4/fp4 post_process + if self.quantization_config and self.quantization_config["quant_method"] in ["compressed-tensors"]: + post_process_weights( + restore_weights_before_load=False, + post_process_quantization=True, + rollout_engines=self.rollout_engines, + ) + ray.get([engine.continue_generation.remote() for engine in self.rollout_engines]) + dist.barrier(group=get_gloo_group()) + + def _sync_weights_to_rollout_engines(self) -> None: use_vllm_packed = self._use_vllm_packed() if use_vllm_packed and self._is_pp_src_rank: logger.info("Using vLLM packed weight sync (bucketed; metadata + trainer_send_weights per bucket)") @@ -350,17 +207,8 @@ def update_weights(self) -> None: if named_tensors: self._update_expert_bucket_weights_from_distributed(named_tensors, pbar=pbar) - dist.barrier(group=get_gloo_group()) - if dist.get_rank() == 0: - # int4/fp4 post_process - if self.quantization_config and self.quantization_config["quant_method"] in ["compressed-tensors"]: - post_process_weights( - restore_weights_before_load=False, - post_process_quantization=True, - rollout_engines=self.rollout_engines, - ) - ray.get([engine.continue_generation.remote() for engine in self.rollout_engines]) - dist.barrier(group=get_gloo_group()) + if self._is_pp_src_rank: + torch.cuda.synchronize() def _use_vllm_packed(self) -> bool: """Use vLLM packed weight transfer (one-shot metadata + trainer_send_weights).""" @@ -523,10 +371,8 @@ def connect_rollout_engines_from_distributed( have heterogeneous TP sizes (e.g. prefill TP=2, decode TP=4), each engine occupies a different number of ranks in the NCCL group. - For vLLM backend, the trainer-side NCCL communicator is created inside a - separate subprocess (_NcclBridge) to avoid conflicts between vLLM's raw - NCCL (PyNcclCommunicator) and the torch.distributed NCCL groups that - Megatron already holds in this process. + Trainer rank 0 uses ``NCCLWeightTransferEngine.trainer_init`` + in-process (StatelessProcessGroup + PyNcclCommunicator). """ if engine_gpu_counts is None: engine_gpu_counts = [args.rollout_num_gpus_per_engine] * len(rollout_engines) @@ -558,18 +404,19 @@ def connect_rollout_engines_from_distributed( device = torch.cuda.current_device() logger.info( - "vLLM weight transfer via NcclBridge: addr=%s port=%d world_size=%d device=%d CVD=%s", + "vLLM in-process weight transfer: addr=%s port=%d world_size=%d device=%d CVD=%s", master_address, master_port, world_size, device, os.environ.get("CUDA_VISIBLE_DEVICES", ""), ) - model_update_groups = _NcclBridge( - master_address=master_address, - master_port=master_port, - world_size=world_size, - device=device, + model_update_groups = NCCLWeightTransferEngine.trainer_init( + { + "master_address": master_address, + "master_port": master_port, + "world_size": world_size, + } ) ray.get(refs) @@ -586,14 +433,12 @@ def disconnect_rollout_engines_from_distributed( Destroy NCCL on training and engines. """ refs = [engine.destroy_weights_update_group.remote(group_name) for engine in rollout_engines] - if isinstance(model_update_groups, _NcclBridge): - model_update_groups.shutdown() ray.get(refs) def update_weights_from_distributed( group_name: str, - group: _NcclBridge, + group: Any, weight_version: int, rollout_engines: Sequence[ActorHandle], converted_named_tensors: Sequence[tuple[str, torch.Tensor]], @@ -601,9 +446,10 @@ def update_weights_from_distributed( packed: bool = False, ) -> list[ObjectRef]: """ - Send metadata (Ray), broadcast tensors (NCCL rank 0 → vLLM engines via - the ``_NcclBridge`` subprocess so that raw NCCL never runs inside the - Megatron trainer process). + Send metadata (Ray), broadcast tensors (NCCL rank 0 → engines). + + The *group* is a vLLM ``PyNcclCommunicator`` from ``trainer_init`` + in the Megatron trainer process. """ kwargs: dict[str, Any] = { "names": [name for name, _ in converted_named_tensors], @@ -616,10 +462,14 @@ def update_weights_from_distributed( refs = [engine.update_weights_from_distributed.remote(**kwargs) for engine in rollout_engines] - if packed: - group.send_weights_packed(list(converted_named_tensors)) - else: - group.broadcast_tensors([param.data for _, param in converted_named_tensors]) + named_gpu_iter = ( + (name, (param.data if hasattr(param, "data") else param).contiguous()) + for name, param in converted_named_tensors + ) + NCCLWeightTransferEngine.trainer_send_weights( + named_gpu_iter, + NCCLTrainerSendWeightsArgs(group=group, packed=packed), + ) return refs diff --git a/slime/backends/vllm_utils/arguments.py b/slime/backends/vllm_utils/arguments.py index 7aa47e258..8669e6548 100644 --- a/slime/backends/vllm_utils/arguments.py +++ b/slime/backends/vllm_utils/arguments.py @@ -217,7 +217,7 @@ def add_vllm_arguments(parser): "--no-vllm-weight-sync-packed", dest="vllm_weight_sync_packed", action="store_false", - help="Disable packed sync; use per-bucket NCCL via NcclBridge instead.", + help="Disable packed sync; send weights per bucket via in-process NCCL (non-packed mode).", ) parser.set_defaults(vllm_weight_sync_packed=True) diff --git a/slime/backends/vllm_utils/vllm_engine.py b/slime/backends/vllm_utils/vllm_engine.py index a598e0441..d50704def 100644 --- a/slime/backends/vllm_utils/vllm_engine.py +++ b/slime/backends/vllm_utils/vllm_engine.py @@ -19,6 +19,8 @@ # vLLM sleep/wake only supports these tags (SGLang also uses ``cuda_graph``, which must be dropped). _VLLM_WAKE_TAGS = frozenset({"weights", "kv_cache"}) +_SKIP_NON_LEADER = {"ok": True, "skipped": True} + def _normalize_vllm_wake_tags(tags: list[str] | None) -> list[str] | None: if not tags: @@ -59,6 +61,17 @@ def _to_local_gpu_id(physical_gpu_id: int) -> int: ) +def _response_json_or_fallback(response: requests.Response) -> dict: + """Parse JSON body; on decode failure return an error-shaped dict (HTTP status already checked).""" + try: + body = response.json() + if isinstance(body, dict): + return body + return {"ok": False, "error": "Response is not a dictionary", "data": body} + except ValueError: + return {"ok": False, "error": "Invalid JSON response", "raw": response.text} + + def _format_v6_uri(addr: str | None) -> str | None: if not addr or addr.startswith("["): return addr @@ -163,7 +176,7 @@ def _forward_vllm_cli_args(args, cmd: list[str]) -> None: serialized = _serialize_for_cli(value) if serialized is None: logger.debug( - "Skipping forward of %s: parsed value %r (%s) cannot be serialized; " "needs vime-side handling.", + "Skipping forward of %s: parsed value %r (%s) cannot be serialized; needs vime-side handling.", vllm_flag, value, type(value).__name__, @@ -413,13 +426,22 @@ def __init__( self.num_gpus_per_engine = num_gpus_per_engine self.process: multiprocessing.Process | None = None self._weight_version: str | None = None - self._is_local_server = not args.rollout_external # Slime runs one vLLM HTTP process per logical engine; multi-node worker rank is not used. self.node_rank = 0 + self.server_host: str | None = None + self.server_port: int | None = None + self._weight_transfer_http_timeout_s: float | None = None def _http_base(self) -> str: + if self.server_host is None or self.server_port is None: + raise RuntimeError("VLLMEngine.init() must be called before HTTP requests") return f"http://{self.server_host}:{self.server_port}" + def _skipped_if_not_leader(self) -> dict | None: + if self.node_rank != 0: + return dict(_SKIP_NON_LEADER) + return None + def init( self, dist_init_addr, @@ -532,29 +554,45 @@ def _post_json(self, endpoint: str, payload: dict, timeout: float) -> requests.R url = f"{self._http_base()}/{endpoint.lstrip('/')}" return requests.post(url, json=payload, timeout=timeout) - def _post_vllm_update_weights_http(self, update_info: dict) -> dict: - """POST ``/update_weights`` with ``{"update_info": ...}`` (vLLM RLHF control plane). - - Same contract as upstream ``examples/online_serving/new_weight_syncing/rlhf_http_nccl.py``: - no ``start_weight_update`` / ``finish_weight_update`` wrapper. - """ - timeout_s = float( - os.environ.get( - "SLIME_VLLM_WEIGHT_TRANSFER_UPDATE_TIMEOUT_SEC", - os.environ.get("SLIME_VLLM_WEIGHT_TRANSFER_HTTP_TIMEOUT_SEC", "900"), + def _weight_transfer_http_timeout(self) -> float: + if self._weight_transfer_http_timeout_s is None: + self._weight_transfer_http_timeout_s = float( + os.environ.get( + "SLIME_VLLM_WEIGHT_TRANSFER_UPDATE_TIMEOUT_SEC", + os.environ.get("SLIME_VLLM_WEIGHT_TRANSFER_HTTP_TIMEOUT_SEC", "900"), + ) ) + return self._weight_transfer_http_timeout_s + + def start_weight_update(self, is_checkpoint_format: bool = True) -> dict: + """``POST /start_weight_update`` (vLLM 0.21+ weight transfer).""" + if skipped := self._skipped_if_not_leader(): + return skipped + response = self._post_json( + "start_weight_update", + {"is_checkpoint_format": is_checkpoint_format}, + timeout=self._weight_transfer_http_timeout(), ) - response = self._post_json("update_weights", {"update_info": update_info}, timeout=timeout_s) response.raise_for_status() - try: - return response.json() - except Exception: - return {"ok": True, "raw": response.text} + return _response_json_or_fallback(response) - def _run_vllm_weight_update(self, update_info: dict, *, is_checkpoint_format: bool = False): - """Backward-compatible alias for non-NCCL ``update_info`` shapes (e.g. tensor/IPC path).""" - del is_checkpoint_format - return self._post_vllm_update_weights_http(update_info) + def finish_weight_update(self) -> dict: + """``POST /finish_weight_update`` (vLLM 0.21+ weight transfer).""" + if skipped := self._skipped_if_not_leader(): + return skipped + response = self._post_json("finish_weight_update", {}, timeout=self._weight_transfer_http_timeout()) + response.raise_for_status() + return _response_json_or_fallback(response) + + def _post_vllm_update_weights_http(self, update_info: dict) -> dict: + """POST ``/update_weights`` with ``{"update_info": ...}`` (vLLM RLHF control plane).""" + response = self._post_json( + "update_weights", + {"update_info": update_info}, + timeout=self._weight_transfer_http_timeout(), + ) + response.raise_for_status() + return _response_json_or_fallback(response) def health_generate(self, timeout: float = 5.0) -> bool: """Return True if ``GET /health`` succeeds (SGLang uses ``GET /health_generate`` for the same role).""" @@ -596,7 +634,7 @@ def update_weights_from_tensor( "format": "serialized_named_tensors", "weight_version": self._weight_version, } - return self._run_vllm_weight_update(update_info, is_checkpoint_format=False) + return self._post_vllm_update_weights_http(update_info) def flush_cache(self): """Clear prefix cache via ``POST /reset_prefix_cache`` (SGLang uses ``GET /flush_cache``).""" @@ -688,10 +726,7 @@ def release_memory_occupation(self): timeout=30, ) response.raise_for_status() - try: - return response.json() - except Exception: - return {"ok": True, "raw": response.text} + return _response_json_or_fallback(response) def resume_memory_occupation(self, tags: list[str] | None = None): """``POST /wake_up`` when sleep mode is on (SGLang: ``POST /resume_memory_occupation``); else a small placeholder dict.""" @@ -707,10 +742,7 @@ def resume_memory_occupation(self, tags: list[str] | None = None): timeout=30, ) response.raise_for_status() - try: - return response.json() - except Exception: - return {"ok": True, "raw": response.text} + return _response_json_or_fallback(response) def check_weights(self, action: str): """No vLLM ``weights_checker`` route; return a placeholder (SGLang posts to ``/weights_checker``).""" @@ -732,16 +764,13 @@ def init_weights_update_group(self, master_address, master_port, rank_offset, wo "world_size": world_size, } } - init_timeout_s = float(os.environ.get("SLIME_VLLM_WEIGHT_TRANSFER_HTTP_TIMEOUT_SEC", "900")) + timeout_s = self._weight_transfer_http_timeout() last_error = None for attempt in range(1, 4): try: - response = self._post_json("init_weight_transfer_engine", payload, timeout=init_timeout_s) + response = self._post_json("init_weight_transfer_engine", payload, timeout=timeout_s) response.raise_for_status() - try: - return response.json() - except Exception: - return {"ok": True, "raw": response.text} + return _response_json_or_fallback(response) except Exception as e: last_error = e if attempt < 3: @@ -796,10 +825,7 @@ def update_weights_from_disk(self, model_path: str, load_format: str | None = No timeout=600, ) response.raise_for_status() - try: - return response.json() - except Exception: - return {"ok": True, "raw": response.text} + return _response_json_or_fallback(response) def pause_generation(self): """``POST /pause`` with mode="keep" (SGLang: ``POST /pause_generation``); returns the ``requests.Response``.""" diff --git a/tests/test_update_weight_from_distributed.py b/tests/test_update_weight_from_distributed.py deleted file mode 100644 index 763992eef..000000000 --- a/tests/test_update_weight_from_distributed.py +++ /dev/null @@ -1,217 +0,0 @@ -"""Unit tests for slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py.""" - -from __future__ import annotations - -import importlib -import inspect -from collections.abc import Iterable -from dataclasses import dataclass, field - -import pytest -import torch - - -@pytest.fixture(scope="module") -def upw(): - return importlib.import_module("slime.backends.megatron_utils.update_weight.update_weight_from_distributed") - - -@dataclass -class _RemoteCall: - args: tuple - kwargs: dict - - -class RecordingRemoteMethod: - def __init__(self, return_value: str = "ref"): - self._return_value = return_value - self.calls: list[_RemoteCall] = [] - - def remote(self, *args, **kwargs): - self.calls.append(_RemoteCall(args=args, kwargs=kwargs)) - return self._return_value - - -@dataclass -class RecordingEngine: - update_weights_from_distributed: RecordingRemoteMethod = field( - default_factory=lambda: RecordingRemoteMethod("ref") - ) - - -@dataclass -class RecordingNcclBridge: - broadcast_calls: list[list[torch.Tensor]] = field(default_factory=list) - packed_calls: list[list[tuple[str, torch.Tensor]]] = field(default_factory=list) - - def broadcast_tensors(self, tensors: Iterable[torch.Tensor]) -> None: - self.broadcast_calls.append(list(tensors)) - - def send_weights_packed(self, named_tensors: Iterable[tuple[str, torch.Tensor]]) -> None: - self.packed_calls.append(list(named_tensors)) - - -def _real_tensors(n: int = 2): - return [(f"layer.{i}.weight", torch.zeros(2, 2)) for i in range(n)] - - -@pytest.mark.unit -def test_signature_no_use_vllm(upw): - sig = inspect.signature(upw.update_weights_from_distributed) - params = sig.parameters - assert "use_vllm" not in params - for p in ("group_name", "group", "weight_version", "rollout_engines", "converted_named_tensors", "packed"): - assert p in params - - -@pytest.mark.unit -def test_signature_rejects_legacy_use_vllm_call(upw): - with pytest.raises(TypeError, match="use_vllm"): - upw.update_weights_from_distributed( - "g", - RecordingNcclBridge(), - 1, - [RecordingEngine()], - _real_tensors(), - use_vllm=True, - packed=False, - ) - - -@pytest.mark.unit -def test_packed_true_dispatches_send_weights_packed(upw): - group = RecordingNcclBridge() - engine = RecordingEngine() - tensors = _real_tensors() - - refs = upw.update_weights_from_distributed("groupA", group, 7, [engine], tensors, packed=True) - - assert len(group.packed_calls) == 1 - assert len(group.broadcast_calls) == 0 - sent = group.packed_calls[0] - assert [n for n, _ in sent] == [n for n, _ in tensors] - assert refs == ["ref"] - - -@pytest.mark.unit -def test_packed_false_dispatches_broadcast_tensors(upw): - group = RecordingNcclBridge() - engine = RecordingEngine() - tensors = _real_tensors() - - refs = upw.update_weights_from_distributed("groupB", group, 7, [engine], tensors, packed=False) - - assert len(group.broadcast_calls) == 1 - assert len(group.packed_calls) == 0 - assert len(group.broadcast_calls[0]) == len(tensors) - assert refs == ["ref"] - - -@pytest.mark.unit -def test_default_packed_is_false(upw): - group = RecordingNcclBridge() - engine = RecordingEngine() - - upw.update_weights_from_distributed("g", group, 1, [engine], _real_tensors()) - - assert len(group.broadcast_calls) == 1 - assert len(group.packed_calls) == 0 - - -@pytest.mark.unit -def test_no_dist_broadcast_fallback(upw, monkeypatch): - import torch.distributed as dist - - seen_broadcast = [] - - def fake_broadcast(*a, **k): - seen_broadcast.append((a, k)) - - monkeypatch.setattr(dist, "broadcast", fake_broadcast) - - group = RecordingNcclBridge() - engine = RecordingEngine() - upw.update_weights_from_distributed("g", group, 1, [engine], _real_tensors(), packed=False) - - assert seen_broadcast == [] - assert group.broadcast_calls - - -@pytest.mark.unit -def test_remote_kwargs_include_packed_true(upw): - group = RecordingNcclBridge() - engine = RecordingEngine() - tensors = _real_tensors(n=1) - - upw.update_weights_from_distributed("myg", group, 42, [engine], tensors, packed=True) - - assert len(engine.update_weights_from_distributed.calls) == 1 - kw = engine.update_weights_from_distributed.calls[0].kwargs - assert kw["packed"] is True - assert kw["group_name"] == "myg" - assert kw["weight_version"] == "42" - assert kw["names"] == ["layer.0.weight"] - assert kw["shapes"] == [torch.Size([2, 2])] - assert kw["dtypes"] == [torch.float32] - - -@pytest.mark.unit -def test_remote_kwargs_include_packed_false(upw): - group = RecordingNcclBridge() - engine = RecordingEngine() - tensors = _real_tensors(n=2) - - upw.update_weights_from_distributed("g", group, 99, [engine], tensors, packed=False) - - kw = engine.update_weights_from_distributed.calls[0].kwargs - assert kw["packed"] is False - assert kw["weight_version"] == "99" - assert kw["names"] == ["layer.0.weight", "layer.1.weight"] - - -@pytest.mark.unit -def test_remote_kwargs_no_use_vllm(upw): - group = RecordingNcclBridge() - engine = RecordingEngine() - - upw.update_weights_from_distributed("g", group, 1, [engine], _real_tensors(), packed=False) - - kw = engine.update_weights_from_distributed.calls[0].kwargs - assert "use_vllm" not in kw - - -@pytest.mark.unit -def test_multiple_engines_each_get_call(upw): - group = RecordingNcclBridge() - engines = [RecordingEngine() for _ in range(3)] - upw.update_weights_from_distributed("g", group, 1, engines, _real_tensors(), packed=True) - for e in engines: - assert len(e.update_weights_from_distributed.calls) == 1 - - -@pytest.mark.unit -def test_empty_tensor_list_still_dispatches(upw): - group = RecordingNcclBridge() - engine = RecordingEngine() - - refs = upw.update_weights_from_distributed("g", group, 1, [engine], [], packed=False) - - assert refs == ["ref"] - kw = engine.update_weights_from_distributed.calls[0].kwargs - assert kw["names"] == [] - assert kw["shapes"] == [] - assert len(group.broadcast_calls) == 1 - assert group.broadcast_calls[0] == [] - - -@pytest.mark.unit -def test_source_no_standalone_use_vllm_param(upw): - src = inspect.getsource(upw) - lines = [line.strip() for line in src.splitlines() if "use_vllm=" in line and "use_vllm_packed" not in line] - assert lines == [] - - -@pytest.mark.unit -def test_source_no_sglang_dist_broadcast_fallback(upw): - fn_src = inspect.getsource(upw.update_weights_from_distributed) - assert "dist.broadcast(" not in fn_src diff --git a/tests/unit/__init__.py b/tests/unit/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/unit/backends/__init__.py b/tests/unit/backends/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/unit/backends/megatron_utils/__init__.py b/tests/unit/backends/megatron_utils/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/unit/backends/megatron_utils/update_weight/__init__.py b/tests/unit/backends/megatron_utils/update_weight/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/unit/backends/megatron_utils/update_weight/test_update_weight_from_distributed.py b/tests/unit/backends/megatron_utils/update_weight/test_update_weight_from_distributed.py new file mode 100644 index 000000000..3eba503f4 --- /dev/null +++ b/tests/unit/backends/megatron_utils/update_weight/test_update_weight_from_distributed.py @@ -0,0 +1,363 @@ +"""Unit tests for slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py.""" + +from __future__ import annotations + +import importlib +import inspect +from dataclasses import dataclass, field + +import pytest +import torch + + +@pytest.fixture(scope="module") +def upw(): + return importlib.import_module("slime.backends.megatron_utils.update_weight.update_weight_from_distributed") + + +@dataclass +class _RemoteCall: + args: tuple + kwargs: dict + + +class RecordingRemoteMethod: + def __init__(self, return_value: str = "ref"): + self._return_value = return_value + self.calls: list[_RemoteCall] = [] + + def remote(self, *args, **kwargs): + self.calls.append(_RemoteCall(args=args, kwargs=kwargs)) + return self._return_value + + +@dataclass +class RecordingEngine: + update_weights_from_distributed: RecordingRemoteMethod = field( + default_factory=lambda: RecordingRemoteMethod("ref") + ) + init_weights_update_group: RecordingRemoteMethod = field(default_factory=lambda: RecordingRemoteMethod("init_ref")) + destroy_weights_update_group: RecordingRemoteMethod = field( + default_factory=lambda: RecordingRemoteMethod("destroy_ref") + ) + start_weight_update: RecordingRemoteMethod = field(default_factory=lambda: RecordingRemoteMethod("start_ref")) + finish_weight_update: RecordingRemoteMethod = field(default_factory=lambda: RecordingRemoteMethod("finish_ref")) + + +@dataclass +class DummyGroup: + token: str = "dummy" + + +def _real_tensors(n: int = 2): + return [(f"layer.{i}.weight", torch.zeros(2, 2)) for i in range(n)] + + +def _make_dummy_nccl_engine(*, send_seen: list[dict] | None = None, init_seen: list[dict] | None = None): + """Build dummy NCCL types; patch on *upw* module (top-level import, not sys.modules).""" + + class DummyNCCLTrainerSendWeightsArgs: + def __init__(self, *, group, packed): + self.group = group + self.packed = packed + + class DummyNCCLWeightTransferEngine: + @staticmethod + def trainer_send_weights(iterator, trainer_args): + if send_seen is not None: + send_seen.append( + { + "items": list(iterator), + "group": trainer_args.group, + "packed": trainer_args.packed, + } + ) + + @staticmethod + def trainer_init(cfg): + if init_seen is not None: + init_seen.append(cfg) + return DummyGroup("group-from-trainer-init") + + return DummyNCCLWeightTransferEngine, DummyNCCLTrainerSendWeightsArgs + + +def _patch_nccl_on_module( + monkeypatch, upw, *, send_seen: list[dict] | None = None, init_seen: list[dict] | None = None +): + dummy_engine, dummy_args = _make_dummy_nccl_engine(send_seen=send_seen, init_seen=init_seen) + monkeypatch.setattr(upw, "NCCLWeightTransferEngine", dummy_engine) + monkeypatch.setattr(upw, "NCCLTrainerSendWeightsArgs", dummy_args) + + +def _patch_trainer_send(monkeypatch, upw, seen: list[dict]) -> None: + _patch_nccl_on_module(monkeypatch, upw, send_seen=seen) + monkeypatch.setattr(upw.torch.cuda, "synchronize", lambda: None) + + +@pytest.mark.unit +def test_signature_no_use_vllm(upw): + sig = inspect.signature(upw.update_weights_from_distributed) + params = sig.parameters + assert "use_vllm" not in params + for p in ("group_name", "group", "weight_version", "rollout_engines", "converted_named_tensors", "packed"): + assert p in params + + +@pytest.mark.unit +def test_signature_rejects_legacy_use_vllm_call(upw): + with pytest.raises(TypeError, match="use_vllm"): + upw.update_weights_from_distributed( + "g", + DummyGroup(), + 1, + [RecordingEngine()], + _real_tensors(), + use_vllm=True, + packed=False, + ) + + +@pytest.mark.unit +def test_packed_true_uses_vllm_trainer_send_weights(upw, monkeypatch): + group = DummyGroup() + engine = RecordingEngine() + tensors = _real_tensors() + seen = [] + _patch_trainer_send(monkeypatch, upw, seen) + + refs = upw.update_weights_from_distributed("groupA", group, 7, [engine], tensors, packed=True) + + assert len(seen) == 1 + sent = seen[0]["items"] + assert [n for n, _ in sent] == [n for n, _ in tensors] + assert seen[0]["group"] is group + assert seen[0]["packed"] is True + assert refs == ["ref"] + + +@pytest.mark.unit +def test_packed_false_still_uses_vllm_trainer_send_weights(upw, monkeypatch): + group = DummyGroup() + engine = RecordingEngine() + tensors = _real_tensors() + seen = [] + _patch_trainer_send(monkeypatch, upw, seen) + + refs = upw.update_weights_from_distributed("groupB", group, 7, [engine], tensors, packed=False) + + assert len(seen) == 1 + assert len(seen[0]["items"]) == len(tensors) + assert seen[0]["packed"] is False + assert refs == ["ref"] + + +@pytest.mark.unit +def test_default_packed_is_false(upw, monkeypatch): + group = DummyGroup() + engine = RecordingEngine() + seen = [] + _patch_trainer_send(monkeypatch, upw, seen) + + upw.update_weights_from_distributed("g", group, 1, [engine], _real_tensors()) + + assert len(seen) == 1 + assert seen[0]["packed"] is False + + +@pytest.mark.unit +def test_no_dist_broadcast_fallback(upw, monkeypatch): + import torch.distributed as dist + + seen_broadcast = [] + seen_send = [] + + def fake_broadcast(*a, **k): + seen_broadcast.append((a, k)) + + monkeypatch.setattr(dist, "broadcast", fake_broadcast) + _patch_trainer_send(monkeypatch, upw, seen_send) + + group = DummyGroup() + engine = RecordingEngine() + upw.update_weights_from_distributed("g", group, 1, [engine], _real_tensors(), packed=False) + + assert seen_broadcast == [] + assert len(seen_send) == 1 + + +@pytest.mark.unit +def test_remote_kwargs_include_packed_true(upw, monkeypatch): + group = DummyGroup() + engine = RecordingEngine() + tensors = _real_tensors(n=1) + seen_send = [] + _patch_trainer_send(monkeypatch, upw, seen_send) + + upw.update_weights_from_distributed("myg", group, 42, [engine], tensors, packed=True) + + assert len(seen_send) == 1 + assert seen_send[0]["packed"] is True + assert len(engine.update_weights_from_distributed.calls) == 1 + kw = engine.update_weights_from_distributed.calls[0].kwargs + assert kw["packed"] is True + assert kw["group_name"] == "myg" + assert kw["weight_version"] == "42" + assert kw["names"] == ["layer.0.weight"] + assert kw["shapes"] == [torch.Size([2, 2])] + assert kw["dtypes"] == [torch.float32] + + +@pytest.mark.unit +def test_remote_kwargs_include_packed_false(upw, monkeypatch): + group = DummyGroup() + engine = RecordingEngine() + tensors = _real_tensors(n=2) + seen_send = [] + _patch_trainer_send(monkeypatch, upw, seen_send) + + upw.update_weights_from_distributed("g", group, 99, [engine], tensors, packed=False) + + assert len(seen_send) == 1 + assert seen_send[0]["packed"] is False + kw = engine.update_weights_from_distributed.calls[0].kwargs + assert kw["packed"] is False + assert kw["weight_version"] == "99" + assert kw["names"] == ["layer.0.weight", "layer.1.weight"] + + +@pytest.mark.unit +def test_remote_kwargs_no_use_vllm(upw, monkeypatch): + group = DummyGroup() + engine = RecordingEngine() + seen_send = [] + _patch_trainer_send(monkeypatch, upw, seen_send) + + upw.update_weights_from_distributed("g", group, 1, [engine], _real_tensors(), packed=False) + + assert len(seen_send) == 1 + kw = engine.update_weights_from_distributed.calls[0].kwargs + assert "use_vllm" not in kw + + +@pytest.mark.unit +def test_multiple_engines_each_get_call(upw, monkeypatch): + group = DummyGroup() + engines = [RecordingEngine() for _ in range(3)] + seen_send = [] + _patch_trainer_send(monkeypatch, upw, seen_send) + + upw.update_weights_from_distributed("g", group, 1, engines, _real_tensors(), packed=True) + assert len(seen_send) == 1 + assert seen_send[0]["packed"] is True + for e in engines: + assert len(e.update_weights_from_distributed.calls) == 1 + + +@pytest.mark.unit +def test_empty_tensor_list_still_dispatches(upw, monkeypatch): + group = DummyGroup() + engine = RecordingEngine() + seen_send = [] + _patch_trainer_send(monkeypatch, upw, seen_send) + + refs = upw.update_weights_from_distributed("g", group, 1, [engine], [], packed=False) + + assert refs == ["ref"] + kw = engine.update_weights_from_distributed.calls[0].kwargs + assert kw["names"] == [] + assert kw["shapes"] == [] + assert len(seen_send) == 1 + assert seen_send[0]["items"] == [] + + +@pytest.mark.unit +def test_source_no_standalone_use_vllm_param(upw): + src = inspect.getsource(upw) + lines = [line.strip() for line in src.splitlines() if "use_vllm=" in line and "use_vllm_packed" not in line] + assert lines == [] + + +@pytest.mark.unit +def test_source_no_sglang_dist_broadcast_fallback(upw): + src = inspect.getsource(upw) + assert "dist.broadcast(" not in src + + +@pytest.mark.unit +def test_source_no_materialized_named_gpu_list(upw): + src = inspect.getsource(upw.update_weights_from_distributed) + assert "named_gpu = []" not in src + assert "named_gpu_iter =" in src + + +@pytest.mark.unit +def test_connect_rollout_engines_always_uses_vllm_trainer_init(upw, monkeypatch): + args = type("Args", (), {"rollout_num_gpus_per_engine": 1})() + engines = [RecordingEngine(), RecordingEngine()] + seen: list[dict] = [] + + _patch_nccl_on_module(monkeypatch, upw, init_seen=seen) + monkeypatch.setattr(upw.torch.cuda, "synchronize", lambda: None) + monkeypatch.setattr(upw.torch.cuda, "empty_cache", lambda: None) + monkeypatch.setattr(upw.ray, "get", lambda refs: refs) + monkeypatch.setattr(upw.ray._private.services, "get_node_ip_address", lambda: "127.0.0.1") + + group = upw.connect_rollout_engines_from_distributed(args, "g", engines, engine_gpu_counts=[1, 2]) + + assert isinstance(group, DummyGroup) + assert len(seen) == 1 + assert seen[0]["master_address"] == "127.0.0.1" + assert seen[0]["world_size"] == 4 # 1 + (1 + 2) + assert len(engines[0].init_weights_update_group.calls) == 1 + assert len(engines[1].init_weights_update_group.calls) == 1 + + +@pytest.mark.unit +def test_weight_update_session_calls_start_and_finish(upw, monkeypatch): + import torch.distributed as dist + + engines = [RecordingEngine(), RecordingEngine()] + ray_refs = [] + barrier_calls: list[object] = [] + + def fake_barrier(*, group=None, **kwargs): + barrier_calls.append(group) + + monkeypatch.setattr(dist, "get_rank", lambda: 0) + monkeypatch.setattr(dist, "barrier", fake_barrier) + monkeypatch.setattr(upw, "get_gloo_group", lambda: "dummy-gloo-group") + monkeypatch.setattr(upw.ray, "get", lambda refs: ray_refs.extend(refs) or refs) + + upw._begin_vllm_weight_update_session(engines) + upw._end_vllm_weight_update_session(engines) + + assert len(engines[0].start_weight_update.calls) == 1 + assert engines[0].start_weight_update.calls[0].kwargs["is_checkpoint_format"] is True + assert len(engines[1].start_weight_update.calls) == 1 + assert len(engines[0].finish_weight_update.calls) == 1 + assert len(engines[1].finish_weight_update.calls) == 1 + assert barrier_calls == ["dummy-gloo-group", "dummy-gloo-group"] + + +@pytest.mark.unit +def test_source_wraps_sync_with_weight_update_session(upw): + src = inspect.getsource(upw.UpdateWeightFromDistributed.update_weights) + assert "_begin_vllm_weight_update_session" in src + assert "_end_vllm_weight_update_session" in src + assert "_sync_weights_to_rollout_engines" in src + + +@pytest.mark.unit +def test_source_uses_nccl_trainer_send_weights_args(upw): + src = inspect.getsource(upw.update_weights_from_distributed) + assert "NCCLTrainerSendWeightsArgs" in src + assert "weight_transfer_compat" not in src + + +@pytest.mark.unit +def test_cuda_sync_once_after_all_buckets_not_per_bucket(upw): + send_src = inspect.getsource(upw.update_weights_from_distributed) + sync_src = inspect.getsource(upw.UpdateWeightFromDistributed._sync_weights_to_rollout_engines) + assert "torch.cuda.synchronize" not in send_src + assert "torch.cuda.synchronize" in sync_src diff --git a/tests/unit/backends/vllm_utils/__init__.py b/tests/unit/backends/vllm_utils/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/unit/backends/vllm_utils/conftest.py b/tests/unit/backends/vllm_utils/conftest.py new file mode 100644 index 000000000..2e71a94b7 --- /dev/null +++ b/tests/unit/backends/vllm_utils/conftest.py @@ -0,0 +1,36 @@ +"""Shared fixtures for vLLM backend unit tests.""" + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest + + +@pytest.fixture +def vllm_args() -> SimpleNamespace: + return SimpleNamespace( + rollout_external=True, + hf_checkpoint="/tmp/model", + router_ip=None, + router_port=None, + num_gpus_per_node=8, + rollout_num_gpus_per_engine=4, + colocate=False, + debug_rollout_only=False, + actor_num_gpus_per_node=4, + actor_num_nodes=1, + use_critic=False, + critic_num_gpus_per_node=0, + critic_num_nodes=0, + ) + + +@pytest.fixture +def vllm_engine(vllm_args): + from slime.backends.vllm_utils.vllm_engine import VLLMEngine + + engine = VLLMEngine(vllm_args, rank=0) + engine.server_host = "127.0.0.1" + engine.server_port = 8765 + return engine diff --git a/tests/test_vllm_arguments.py b/tests/unit/backends/vllm_utils/test_arguments.py similarity index 99% rename from tests/test_vllm_arguments.py rename to tests/unit/backends/vllm_utils/test_arguments.py index 9a0067799..89c50e224 100644 --- a/tests/test_vllm_arguments.py +++ b/tests/unit/backends/vllm_utils/test_arguments.py @@ -1,4 +1,4 @@ -"""Unit tests for slime/backends/vllm_utils/arguments.py.""" +"""Unit tests for ``slime.backends.vllm_utils.arguments``.""" from __future__ import annotations diff --git a/tests/unit/backends/vllm_utils/test_vllm_engine.py b/tests/unit/backends/vllm_utils/test_vllm_engine.py new file mode 100644 index 000000000..acda0879c --- /dev/null +++ b/tests/unit/backends/vllm_utils/test_vllm_engine.py @@ -0,0 +1,341 @@ +"""Unit tests for ``slime.backends.vllm_utils.vllm_engine``.""" + +from __future__ import annotations + +import dataclasses +import json + +import pytest +import requests +import torch + +from slime.backends.vllm_utils import vllm_engine as mod + + +class _MockResponse: + def __init__(self, *, json_data: dict | None = None, text: str = "", status_code: int = 200): + self._json_data = json_data + self.text = text + self.status_code = status_code + + def raise_for_status(self) -> None: + if self.status_code >= 400: + raise RuntimeError(f"HTTP {self.status_code}") + + def json(self) -> dict: + if self._json_data is None: + raise ValueError("no json") + return self._json_data + + +@pytest.mark.unit +def test_normalize_vllm_wake_tags_drops_unsupported(): + assert mod._normalize_vllm_wake_tags(["weights", "cuda_graph", "kv_cache"]) == ["weights", "kv_cache"] + + +@pytest.mark.unit +def test_normalize_vllm_wake_tags_empty_becomes_none(): + assert mod._normalize_vllm_wake_tags(["cuda_graph"]) is None + + +@pytest.mark.unit +def test_format_v6_uri_wraps_ipv6(): + assert mod._format_v6_uri("2001:db8::1") == "[2001:db8::1]" + + +@pytest.mark.unit +def test_format_v6_uri_ipv4_unchanged(): + assert mod._format_v6_uri("10.0.0.1") == "10.0.0.1" + + +@pytest.mark.unit +def test_to_local_gpu_id_without_cvd(): + assert mod._to_local_gpu_id(3) == 3 + + +@pytest.mark.unit +def test_to_local_gpu_id_maps_physical_id(monkeypatch): + monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "2,3,4") + assert mod._to_local_gpu_id(3) == 1 + + +@pytest.mark.unit +def test_get_base_gpu_id_colocate(vllm_args): + vllm_args.colocate = True + vllm_args.num_gpus_per_node = 8 + vllm_args.rollout_num_gpus_per_engine = 4 + assert mod.get_base_gpu_id(vllm_args, rank=1) == 4 + + +@pytest.mark.unit +def test_start_weight_update_posts_four_phase_endpoint(vllm_engine, monkeypatch): + calls: list[tuple] = [] + + def fake_post(endpoint: str, payload: dict, timeout: float): + calls.append((endpoint, payload, timeout)) + return _MockResponse(json_data={"ok": True}) + + monkeypatch.setattr(vllm_engine, "_post_json", fake_post) + + result = vllm_engine.start_weight_update(is_checkpoint_format=True) + + assert result == {"ok": True} + assert len(calls) == 1 + assert calls[0][0] == "start_weight_update" + assert calls[0][1] == {"is_checkpoint_format": True} + assert calls[0][2] == vllm_engine._weight_transfer_http_timeout() + + +@pytest.mark.unit +def test_finish_weight_update_posts_empty_body(vllm_engine, monkeypatch): + calls: list[tuple] = [] + + def fake_post(endpoint: str, payload: dict, timeout: float): + calls.append((endpoint, payload, timeout)) + return _MockResponse(json_data={"done": True}) + + monkeypatch.setattr(vllm_engine, "_post_json", fake_post) + + result = vllm_engine.finish_weight_update() + + assert result == {"done": True} + assert calls == [("finish_weight_update", {}, vllm_engine._weight_transfer_http_timeout())] + + +@pytest.mark.unit +def test_start_weight_update_skipped_when_node_rank_nonzero(vllm_engine, monkeypatch): + vllm_engine.node_rank = 1 + monkeypatch.setattr( + vllm_engine, + "_post_json", + lambda *a, **k: pytest.fail("should not POST"), + ) + + assert vllm_engine.start_weight_update() == {"ok": True, "skipped": True} + + +@pytest.mark.unit +def test_update_weights_from_distributed_posts_update_weights_without_checkpoint_flag(vllm_engine, monkeypatch): + calls: list[dict] = [] + + def fake_post_vllm(update_info: dict) -> dict: + calls.append(update_info) + return {"ok": True} + + monkeypatch.setattr(vllm_engine, "_post_vllm_update_weights_http", fake_post_vllm) + + names = ["layer.0.weight"] + dtypes = [torch.float32] + shapes = [torch.Size([2, 2])] + + vllm_engine.update_weights_from_distributed( + names, + dtypes, + shapes, + group_name="slime-pp_0", + weight_version="7", + packed=True, + ) + + assert len(calls) == 1 + info = calls[0] + assert info["names"] == names + assert info["dtype_names"] == ["float32"] + assert info["shapes"] == [[2, 2]] + assert info["packed"] is True + assert "is_checkpoint_format" not in info + assert vllm_engine._weight_version == "7" + + +@pytest.mark.unit +def test_post_vllm_update_weights_http_wraps_update_info(vllm_engine, monkeypatch): + seen: list[tuple] = [] + + def fake_post(endpoint: str, payload: dict, timeout: float): + seen.append((endpoint, payload, timeout)) + return _MockResponse(json_data={"status": "ok"}) + + monkeypatch.setattr(vllm_engine, "_post_json", fake_post) + + result = vllm_engine._post_vllm_update_weights_http({"names": ["w"], "packed": False}) + + assert result == {"status": "ok"} + assert seen[0][0] == "update_weights" + assert seen[0][1] == {"update_info": {"names": ["w"], "packed": False}} + + +@pytest.mark.unit +def test_weight_transfer_http_timeout_reads_env(vllm_engine, monkeypatch): + monkeypatch.setenv("SLIME_VLLM_WEIGHT_TRANSFER_UPDATE_TIMEOUT_SEC", "123.5") + assert vllm_engine._weight_transfer_http_timeout() == 123.5 + + +@pytest.mark.unit +def test_weight_transfer_http_timeout_fallback_to_legacy_env(vllm_engine, monkeypatch): + monkeypatch.delenv("SLIME_VLLM_WEIGHT_TRANSFER_UPDATE_TIMEOUT_SEC", raising=False) + monkeypatch.setenv("SLIME_VLLM_WEIGHT_TRANSFER_HTTP_TIMEOUT_SEC", "42") + assert vllm_engine._weight_transfer_http_timeout() == 42.0 + + +@pytest.mark.unit +def test_response_json_or_fallback_parses_dict(): + response = _MockResponse(json_data={"status": "ready"}) + assert mod._response_json_or_fallback(response) == {"status": "ready"} + + +@pytest.mark.unit +def test_response_json_or_fallback_non_dict_wrapped(): + response = _MockResponse() + response.json = lambda: ["a", "b"] # type: ignore[method-assign] + assert mod._response_json_or_fallback(response) == { + "ok": False, + "error": "Response is not a dictionary", + "data": ["a", "b"], + } + + +@pytest.mark.unit +def test_response_json_or_fallback_invalid_json(): + response = _MockResponse(text="not-json") + response.json = lambda: (_ for _ in ()).throw(ValueError("no json")) # type: ignore[method-assign] + assert mod._response_json_or_fallback(response) == { + "ok": False, + "error": "Invalid JSON response", + "raw": "not-json", + } + + +@pytest.mark.unit +def test_http_base_requires_init(vllm_args): + from slime.backends.vllm_utils.vllm_engine import VLLMEngine + + engine = VLLMEngine(vllm_args, rank=0) + with pytest.raises(RuntimeError, match="init\\(\\)"): + engine._http_base() + + +@pytest.mark.unit +def test_http_base_ipv6_host(vllm_engine): + vllm_engine.server_host = "[2001:db8::1]" + assert vllm_engine._http_base() == "http://[2001:db8::1]:8765" + + +@pytest.mark.unit +def test_redact_cmd_for_log_masks_hf_token(): + cmd = ["vllm", "serve", "model", "--hf-token", "secret-token", "--port", "8000"] + logged = mod._redact_cmd_for_log(cmd) + assert "secret-token" not in logged + assert "***" in logged + + +@pytest.mark.unit +def test_serialize_for_cli_primitives(): + assert mod._serialize_for_cli(42) == "42" + assert mod._serialize_for_cli(True) == "True" + assert mod._serialize_for_cli({"backend": "nccl"}) == json.dumps({"backend": "nccl"}) + + +@pytest.mark.unit +def test_serialize_for_cli_dataclass(): + @dataclasses.dataclass + class _Cfg: + backend: str = "nccl" + + out = mod._serialize_for_cli(_Cfg()) + assert json.loads(out) == {"backend": "nccl"} + + +@pytest.mark.unit +def test_get_base_gpu_id_with_critic_offset(vllm_args): + vllm_args.colocate = False + vllm_args.debug_rollout_only = False + vllm_args.actor_num_gpus_per_node = 4 + vllm_args.actor_num_nodes = 1 + vllm_args.use_critic = True + vllm_args.critic_num_gpus_per_node = 2 + vllm_args.critic_num_nodes = 1 + vllm_args.num_gpus_per_node = 8 + vllm_args.rollout_num_gpus_per_engine = 2 + # actor 4 + critic 2 + rank0*2 = 6 + assert mod.get_base_gpu_id(vllm_args, rank=0) == 6 + + +@pytest.mark.unit +def test_resume_memory_occupation_wake_tags_query(vllm_engine, monkeypatch): + seen: list[tuple] = [] + + def fake_post(url, *, params=None, timeout=30, json=None): + seen.append((url, params, timeout, json)) + return _MockResponse(json_data={"ok": True}) + + vllm_args = vllm_engine.args + vllm_args.vllm_enable_sleep_mode = True + monkeypatch.setattr(mod.requests, "post", fake_post) + + vllm_engine.resume_memory_occupation(tags=["weights", "cuda_graph"]) + + assert len(seen) == 1 + assert seen[0][1] == [("tags", "weights")] + + +@pytest.mark.unit +def test_resume_memory_occupation_skips_when_sleep_disabled(vllm_engine): + vllm_engine.args.vllm_enable_sleep_mode = False + assert vllm_engine.resume_memory_occupation() == {"ok": True, "sleep_mode": False} + + +@pytest.mark.unit +def test_init_weights_update_group_retries_then_succeeds(vllm_engine, monkeypatch): + attempts = {"n": 0} + + def fake_post(endpoint: str, payload: dict, timeout: float): + attempts["n"] += 1 + if attempts["n"] < 2: + raise requests.ConnectionError("transient") + return _MockResponse(json_data={"initialized": True}) + + monkeypatch.setattr(vllm_engine, "_post_json", fake_post) + monkeypatch.setattr(mod.time, "sleep", lambda *_a, **_k: None) + + result = vllm_engine.init_weights_update_group( + "127.0.0.1", + 29500, + rank_offset=1, + world_size=4, + group_name="unused", + backend="nccl", + ) + + assert result == {"initialized": True} + assert attempts["n"] == 2 + + +@pytest.mark.unit +def test_init_weights_update_group_raises_after_three_failures(vllm_engine, monkeypatch): + monkeypatch.setattr( + vllm_engine, + "_post_json", + lambda *a, **k: (_ for _ in ()).throw(requests.ConnectionError("down")), + ) + monkeypatch.setattr(mod.time, "sleep", lambda *_a, **_k: None) + + with pytest.raises(RuntimeError, match="init_weight_transfer_engine failed"): + vllm_engine.init_weights_update_group( + "127.0.0.1", + 29500, + rank_offset=1, + world_size=4, + group_name="g", + backend="nccl", + ) + + +@pytest.mark.unit +def test_finish_weight_update_skipped_when_node_rank_nonzero(vllm_engine, monkeypatch): + vllm_engine.node_rank = 1 + monkeypatch.setattr( + vllm_engine, + "_post_json", + lambda *a, **k: pytest.fail("should not POST"), + ) + assert vllm_engine.finish_weight_update() == {"ok": True, "skipped": True} diff --git a/tests/unit/conftest.py b/tests/unit/conftest.py new file mode 100644 index 000000000..b5df2b783 --- /dev/null +++ b/tests/unit/conftest.py @@ -0,0 +1,19 @@ +"""Unit-test collection hooks (e.g. stub optional heavy deps on dev machines).""" + +from __future__ import annotations + +import sys +from unittest.mock import MagicMock + + +def _ensure_ray_stub() -> None: + if "ray" in sys.modules: + return + ray = MagicMock() + sys.modules["ray"] = ray + sys.modules["ray._private"] = MagicMock() + sys.modules["ray._private.services"] = MagicMock() + sys.modules["ray.actor"] = MagicMock() + + +_ensure_ray_stub()