diff --git a/python/sglang/srt/debug_utils/dumper.py b/python/sglang/srt/debug_utils/dumper.py index 9a9020add..a3fb617b6 100644 --- a/python/sglang/srt/debug_utils/dumper.py +++ b/python/sglang/srt/debug_utils/dumper.py @@ -140,12 +140,49 @@ class DumperConfig(_BaseConfig): server_port: str = "-1" non_intrusive_mode: str = "core" source_patcher_config: Optional[str] = None + grafter_enable: bool = False + grafter_role: str = "" # required if enabled: "baseline" or "target" + grafter_b2t_filter: Optional[str] = None # names flowing baseline -> target + grafter_master_address: str = "" # required if enabled + grafter_master_port: int = -1 # required if enabled (positive port) + grafter_baseline_world_size: int = -1 # required if enabled + grafter_target_world_size: int = -1 # required if enabled + grafter_backend: str = "nccl" + grafter_group_name: str = "graft" + grafter_timeout: int = 300 @classmethod def _env_prefix(cls) -> str: # NOTE: should not be `SGLANG_DUMPER_`, otherwise it is weird when dumping Megatron in Miles return "DUMPER_" + def __post_init__(self) -> None: + super().__post_init__() + if self.grafter_enable: + assert self.grafter_role in ("baseline", "target"), ( + f"grafter_role must be 'baseline' or 'target' when grafter_enable=True, " + f"got {self.grafter_role!r}" + ) + assert ( + self.grafter_master_address + ), "grafter_master_address must be set when grafter_enable=True" + assert self.grafter_master_port > 0, ( + f"grafter_master_port must be a positive port when grafter_enable=True, " + f"got {self.grafter_master_port}" + ) + assert self.grafter_baseline_world_size > 0, ( + f"grafter_baseline_world_size must be > 0 when grafter_enable=True, " + f"got {self.grafter_baseline_world_size}" + ) + assert self.grafter_target_world_size > 0, ( + f"grafter_target_world_size must be > 0 when grafter_enable=True, " + f"got {self.grafter_target_world_size}" + ) + assert self.grafter_b2t_filter is not None, ( + "grafter_enable=True but grafter_b2t_filter is not set; " + "nothing would ever be grafted" + ) + @property def server_port_parsed(self) -> Optional[Union[int, Literal["reuse"]]]: raw = self.server_port @@ -203,6 +240,7 @@ class _Dumper: self._config = config self._state = _DumperState() self._non_intrusives: list["_NonIntrusiveDumper"] = [] + self._grafter = _Grafter(config=config) # ------------------------------- public :: core --------------------------------- @@ -432,6 +470,7 @@ class _Dumper: recompute_meta = recompute_status.to_pseudo_parallel_meta() value = _materialize_value(value) + self._grafter.maybe_intercept(value=value, tags=tags) if enable_value: self._dump_single( @@ -742,6 +781,102 @@ def _register_forward_hook_or_replace_fn( raise ValueError(f"Unknown mode {mode!r}") +# -------------------------------------- grafter ------------------------------------------ + + +class _GraftRole(enum.Enum): + BASELINE = "baseline" + TARGET = "target" + + +class _Grafter: + """1+1 cross-system tensor grafter. + + Both sides set the SAME `grafter_b2t_filter` (names that flow + baseline -> target). The only per-side difference is `grafter_role`, + which tells the side whether it's the sender (baseline) or the + receiver (target). Receiver overwrites its local target tensor with + the sender's via `value.copy_()`. + """ + + def __init__(self, *, config: DumperConfig) -> None: + self._config = config + self._pg: Optional[dist.ProcessGroup] = None + + def maybe_intercept(self, *, value, tags: dict) -> None: + cfg = self._config + if not cfg.grafter_enable: + return + + if not self._match(cfg.grafter_b2t_filter, tags): + return + + if not isinstance(value, torch.Tensor): + _log( + f"[Grafter] tags={tags} matched grafter_b2t_filter but " + f"value is not a torch.Tensor (got type={type(value).__name__}); " + f"skipping graft. Common cause: dumper.dump called with a " + f"non-tensor value (dict, list, ...) on this name. Either " + f"narrow the filter or wrap the value in a tensor." + ) + return + + self._ensure_group() + role = _GraftRole(cfg.grafter_role) + + # b2t with 1+1: baseline rank is sender (graft rank 0), target is recv + # (graft rank 1). Use broadcast_object_list so receiver can have an + # arbitrarily shaped placeholder; sender ships a pickled tensor. + obj_list: list = [None] + if role == _GraftRole.BASELINE: + obj_list = [value] + _log(f"[Grafter] send role=baseline tags={tags}") + dist.broadcast_object_list(obj_list, src=0, group=self._pg) + if role == _GraftRole.TARGET: + received = obj_list[0] + if isinstance(received, torch.Tensor): + # Pickled CUDA tensors restore to their original-device name; + # that may not match this process's local device, so normalize. + received = received.to(value.device) + _log(f"[Grafter] recv role=target tags={tags}") + value.copy_(received) + + @staticmethod + def _match(expr: Optional[str], tags: dict) -> bool: + if expr is None: + return False + return _evaluate_filter(expr, tags) + + def _ensure_group(self) -> None: + if self._pg is not None: + return + + cfg = self._config + assert ( + dist.is_initialized() + ), "[Grafter] default torch.distributed must be initialized" + role = _GraftRole(cfg.grafter_role) + global_rank = 0 if role == _GraftRole.BASELINE else 1 + total_world = cfg.grafter_baseline_world_size + cfg.grafter_target_world_size + init_method = f"tcp://{cfg.grafter_master_address}:{cfg.grafter_master_port}" + _log( + f"[Grafter] init group: role={role.value} " + f"rank={global_rank} init_method={init_method} " + f"backend={cfg.grafter_backend} name={cfg.grafter_group_name}" + ) + self._pg = _collective_with_timeout( + lambda: _init_custom_process_group( + backend=cfg.grafter_backend, + init_method=init_method, + world_size=total_world, + rank=global_rank, + group_name=cfg.grafter_group_name, + ), + operation_name="_init_custom_process_group in _Grafter", + timeout_seconds=cfg.grafter_timeout, + ) + + # -------------------------------------- util fn ------------------------------------------ @@ -1172,6 +1307,61 @@ def _get_local_ip_by_remote() -> Optional[str]: return None +def _init_custom_process_group( + *, + backend: str, + init_method: str, + world_size: int, + rank: int, + group_name: str, + timeout=None, +): + """Build a fresh torch.distributed process group, separate from the default + one and any other custom groups (e.g. RLHF weight-update groups). Used by + the grafter to bridge baseline and target systems. + + Adapted from sglang.srt.utils.common.init_custom_process_group; inlined + here to keep dumper.py free of cross-file imports. + """ + from torch.distributed.distributed_c10d import ( + Backend, + PrefixStore, + _new_process_group_helper, + _world, + default_pg_timeout, + rendezvous, + ) + + if timeout is None: + timeout = default_pg_timeout + + rendezvous_iterator = rendezvous(init_method, rank, world_size, timeout=timeout) + store, rank, world_size = next(rendezvous_iterator) + store.set_timeout(timeout) + store = PrefixStore(group_name, store) + + backend_obj = Backend(backend) + # PyTorch 2.6 renamed `pg_options` to `backend_options`. + torch_major_minor = tuple( + int(x) for x in torch.__version__.split("+")[0].split(".")[:2] + ) + pg_options_param_name = ( + "backend_options" if torch_major_minor >= (2, 6) else "pg_options" + ) + pg, _ = _new_process_group_helper( + world_size, + rank, + [], + backend_obj, + store, + group_name=group_name, + **{pg_options_param_name: None}, + timeout=timeout, + ) + _world.pg_group_ranks[pg] = {i: i for i in range(world_size)} + return pg + + # -------------------------------------- framework plugins ------------------------------------------ diff --git a/test/registered/debug_utils/test_dumper.py b/test/registered/debug_utils/test_dumper.py index a180fe9a7..bb293f3db 100644 --- a/test/registered/debug_utils/test_dumper.py +++ b/test/registered/debug_utils/test_dumper.py @@ -4,8 +4,10 @@ import os import sys import threading import time +import traceback from contextlib import contextmanager from pathlib import Path +from typing import Optional import pytest import requests @@ -20,6 +22,7 @@ from sglang.srt.debug_utils.dumper import ( _Dumper, _format_tags, _get_default_exp_name, + _Grafter, _log, _map_tensor, _materialize_value, @@ -2592,5 +2595,312 @@ class TestRecomputeStatus: assert _detect_recompute_status() == _RecomputeStatus.DISABLED +class TestGrafterConfig: + def test_from_env_role(self): + with temp_set_env( + DUMPER_GRAFTER_ENABLE="1", + DUMPER_GRAFTER_ROLE="baseline", + DUMPER_GRAFTER_MASTER_ADDRESS="127.0.0.1", + DUMPER_GRAFTER_MASTER_PORT="29500", + DUMPER_GRAFTER_BASELINE_WORLD_SIZE="1", + DUMPER_GRAFTER_TARGET_WORLD_SIZE="1", + DUMPER_GRAFTER_B2T_FILTER="name == 'x'", + ): + cfg = DumperConfig.from_env() + assert cfg.grafter_enable is True + assert cfg.grafter_role == "baseline" + assert cfg.grafter_b2t_filter == "name == 'x'" + assert cfg.grafter_master_port == 29500 + + def test_enable_without_required_fields_raises(self): + # missing role + with pytest.raises(AssertionError, match="grafter_role"): + DumperConfig(grafter_enable=True) + # missing master_address + with pytest.raises(AssertionError, match="grafter_master_address"): + DumperConfig(grafter_enable=True, grafter_role="baseline") + # non-positive port + with pytest.raises(AssertionError, match="grafter_master_port"): + DumperConfig( + grafter_enable=True, + grafter_role="baseline", + grafter_master_address="127.0.0.1", + ) + # missing baseline_world_size + with pytest.raises(AssertionError, match="grafter_baseline_world_size"): + DumperConfig( + grafter_enable=True, + grafter_role="baseline", + grafter_master_address="127.0.0.1", + grafter_master_port=29500, + ) + # missing target_world_size + with pytest.raises(AssertionError, match="grafter_target_world_size"): + DumperConfig( + grafter_enable=True, + grafter_role="baseline", + grafter_master_address="127.0.0.1", + grafter_master_port=29500, + grafter_baseline_world_size=1, + ) + # no filter set + with pytest.raises(AssertionError, match="grafter_b2t_filter"): + DumperConfig( + grafter_enable=True, + grafter_role="baseline", + grafter_master_address="127.0.0.1", + grafter_master_port=29500, + grafter_baseline_world_size=1, + grafter_target_world_size=1, + ) + + def test_disabled_does_not_validate_other_fields(self): + # All grafter_* fields can be left at their absurd defaults when + # grafter_enable is False. + cfg = DumperConfig(grafter_enable=False) + assert cfg.grafter_enable is False + + +def _unit_grafter_config(**overrides) -> DumperConfig: + """Build a fully-valid DumperConfig for unit tests of `_Grafter` filter + matching, without spinning up a process group. Dummy values for required + fields are never reached because these tests short-circuit before + `_ensure_group` runs. + """ + base = dict( + grafter_enable=True, + grafter_role="baseline", + grafter_master_address="127.0.0.1", + grafter_master_port=12345, + grafter_baseline_world_size=1, + grafter_target_world_size=1, + grafter_b2t_filter="name == 'x'", + ) + base.update(overrides) + return DumperConfig(**base) + + +class TestGrafterFilterMatching: + """Unit tests for the filter-matching short-circuit logic.""" + + def test_disabled_returns_silently(self): + grafter = _Grafter(config=_unit_grafter_config(grafter_enable=False)) + grafter.maybe_intercept(value=torch.zeros(2), tags={"name": "x"}) + assert grafter._pg is None # never initialized + + def test_unmatched_name_returns_silently(self): + grafter = _Grafter(config=_unit_grafter_config()) + grafter.maybe_intercept(value=torch.zeros(2), tags={"name": "z"}) + assert grafter._pg is None + + def test_matched_non_tensor_prints_and_skips(self): + """Non-tensor that matches a filter -> log explanation, then skip.""" + grafter = _Grafter(config=_unit_grafter_config()) + with _capture_stdout() as captured: + grafter.maybe_intercept(value={"not": "a tensor"}, tags={"name": "x"}) + out = captured.getvalue() + assert grafter._pg is None + assert "value is not a torch.Tensor" in out, out + + +def _run_graft_test(worker_func, **kwargs): + """Spawn one GPU-using process per role (rank 0 = baseline, rank 1 = target). + + Limited to 1+1 because the CI fleet has only 2 GPUs. Each process + initializes its OWN default PG (nccl, world_size=1) from the start, + mirroring production where baseline and target are independently launched. + """ + import torch.multiprocessing as mp + + role_ports = [find_available_port(29700 + i * 100) for i in range(2)] + + ctx = mp.get_context("spawn") + result_queue = ctx.Queue() + processes = [] + for rank in range(2): + p = ctx.Process( + target=_graft_worker_entry, + args=(rank, role_ports[rank], worker_func, result_queue, kwargs), + ) + p.start() + processes.append(p) + + for p in processes: + p.join() + + errors = [result_queue.get() for _ in range(2)] + errors = [e for e in errors if e] + if errors: + raise AssertionError("\n".join(errors)) + + +def _graft_worker_entry(rank, role_port, worker_func, result_queue, kwargs): + torch.cuda.set_device(rank) + dist.init_process_group( + backend="nccl", + init_method=f"tcp://127.0.0.1:{role_port}", + world_size=1, + rank=0, + ) + try: + worker_func(rank=rank, **kwargs) + result_queue.put(None) + except Exception as e: + result_queue.put(f"rank={rank}: {e}\n{traceback.format_exc()}") + finally: + dist.destroy_process_group() + + +def _make_grafter_test_config( + *, + rank: int, + graft_port: int, + group_name: str, + timeout: int = 30, + b2t_filter: Optional[str] = "name == 'x'", +) -> DumperConfig: + """Build a DumperConfig for distributed grafter tests. rank 0 -> baseline, + rank 1 -> target. Both sides are world_size=1 within their own role's + default PG. + """ + role = "baseline" if rank == 0 else "target" + return DumperConfig( + grafter_enable=True, + grafter_role=role, + grafter_b2t_filter=b2t_filter, + grafter_master_address="127.0.0.1", + grafter_master_port=graft_port, + grafter_baseline_world_size=1, + grafter_target_world_size=1, + grafter_group_name=group_name, + grafter_timeout=timeout, + ) + + +class TestGrafterDistributed: + def test_b2t_copy_roundtrip(self): + """Baseline (rank 0) sends 'x' to target (rank 1), target.copy_'s it.""" + graft_port = find_available_port(29600) + _run_graft_test( + self._test_b2t_func, graft_port=graft_port, group_name="grafter_b2t" + ) + + @staticmethod + def _test_b2t_func(rank, graft_port, group_name): + grafter = _Grafter( + config=_make_grafter_test_config( + rank=rank, graft_port=graft_port, group_name=group_name + ) + ) + try: + if rank == 0: + tensor = torch.tensor([1.0, 2.0, 3.0], device="cuda:0") + grafter.maybe_intercept(value=tensor, tags={"name": "x"}) + else: + target = torch.zeros(3, device="cuda:1") + grafter.maybe_intercept(value=target, tags={"name": "x"}) + assert target.tolist() == [1.0, 2.0, 3.0], f"got {target.tolist()}" + finally: + if grafter._pg is not None: + dist.destroy_process_group(grafter._pg) + + def test_unmatched_name_skipped(self): + graft_port = find_available_port(29620) + _run_graft_test( + self._test_unmatched_func, + graft_port=graft_port, + group_name="grafter_unmatched", + ) + + @staticmethod + def _test_unmatched_func(rank, graft_port, group_name): + grafter = _Grafter( + config=_make_grafter_test_config( + rank=rank, graft_port=graft_port, group_name=group_name + ) + ) + try: + target = torch.tensor([7.0, 7.0, 7.0], device=f"cuda:{rank}") + grafter.maybe_intercept(value=target, tags={"name": "other"}) + assert target.tolist() == [7.0, 7.0, 7.0], "tensor must not be modified" + assert grafter._pg is None, "group must not init for unmatched name" + finally: + if grafter._pg is not None: + dist.destroy_process_group(grafter._pg) + + def test_init_timeout_warns(self): + graft_port = find_available_port(29630) + _run_graft_test( + self._test_init_timeout_func, + graft_port=graft_port, + group_name="grafter_timeout", + ) + + @staticmethod + def _test_init_timeout_func(rank, graft_port, group_name): + grafter = _Grafter( + config=_make_grafter_test_config( + rank=rank, graft_port=graft_port, group_name=group_name, timeout=2 + ) + ) + try: + with _capture_stdout() as captured: + if rank == 1: + time.sleep(4) + tensor = torch.tensor([1.0, 2.0, 3.0], device=f"cuda:{rank}") + if rank == 0: + grafter.maybe_intercept(value=tensor, tags={"name": "x"}) + else: + target = torch.zeros(3, device=f"cuda:{rank}") + grafter.maybe_intercept(value=target, tags={"name": "x"}) + output = captured.getvalue() + if rank == 0: + assert "WARNING" in output, output + assert "has not completed after 2s" in output, output + finally: + if grafter._pg is not None: + dist.destroy_process_group(grafter._pg) + + def test_group_init_is_cached_across_calls(self): + """`_ensure_group` runs lazily on first matched dump and caches `_pg`; + subsequent matched dumps must reuse the same object.""" + graft_port = find_available_port(29660) + _run_graft_test( + self._test_group_cache_func, + graft_port=graft_port, + group_name="grafter_cache", + ) + + @staticmethod + def _test_group_cache_func(rank, graft_port, group_name): + grafter = _Grafter( + config=_make_grafter_test_config( + rank=rank, graft_port=graft_port, group_name=group_name + ) + ) + try: + if rank == 0: + t1 = torch.tensor([1.0, 2.0, 3.0], device="cuda:0") + t2 = torch.tensor([4.0, 5.0, 6.0], device="cuda:0") + grafter.maybe_intercept(value=t1, tags={"name": "x"}) + pg_after_first = grafter._pg + assert pg_after_first is not None + grafter.maybe_intercept(value=t2, tags={"name": "x"}) + assert grafter._pg is pg_after_first, "_pg must be cached" + else: + t1 = torch.zeros(3, device="cuda:1") + t2 = torch.zeros(3, device="cuda:1") + grafter.maybe_intercept(value=t1, tags={"name": "x"}) + pg_after_first = grafter._pg + assert pg_after_first is not None + grafter.maybe_intercept(value=t2, tags={"name": "x"}) + assert grafter._pg is pg_after_first + assert t1.tolist() == [1.0, 2.0, 3.0] + assert t2.tolist() == [4.0, 5.0, 6.0] + finally: + if grafter._pg is not None: + dist.destroy_process_group(grafter._pg) + + if __name__ == "__main__": sys.exit(pytest.main([__file__]))