Support cross-system tensor grafting in dumper (#24507)

This commit is contained in:
fzyzcjy
2026-05-06 16:55:40 +08:00
committed by GitHub
parent 61104d7d0a
commit 58487e68e5
2 changed files with 500 additions and 0 deletions
+190
View File
@@ -140,12 +140,49 @@ class DumperConfig(_BaseConfig):
server_port: str = "-1" server_port: str = "-1"
non_intrusive_mode: str = "core" non_intrusive_mode: str = "core"
source_patcher_config: Optional[str] = None 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 @classmethod
def _env_prefix(cls) -> str: def _env_prefix(cls) -> str:
# NOTE: should not be `SGLANG_DUMPER_`, otherwise it is weird when dumping Megatron in Miles # NOTE: should not be `SGLANG_DUMPER_`, otherwise it is weird when dumping Megatron in Miles
return "DUMPER_" 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 @property
def server_port_parsed(self) -> Optional[Union[int, Literal["reuse"]]]: def server_port_parsed(self) -> Optional[Union[int, Literal["reuse"]]]:
raw = self.server_port raw = self.server_port
@@ -203,6 +240,7 @@ class _Dumper:
self._config = config self._config = config
self._state = _DumperState() self._state = _DumperState()
self._non_intrusives: list["_NonIntrusiveDumper"] = [] self._non_intrusives: list["_NonIntrusiveDumper"] = []
self._grafter = _Grafter(config=config)
# ------------------------------- public :: core --------------------------------- # ------------------------------- public :: core ---------------------------------
@@ -432,6 +470,7 @@ class _Dumper:
recompute_meta = recompute_status.to_pseudo_parallel_meta() recompute_meta = recompute_status.to_pseudo_parallel_meta()
value = _materialize_value(value) value = _materialize_value(value)
self._grafter.maybe_intercept(value=value, tags=tags)
if enable_value: if enable_value:
self._dump_single( self._dump_single(
@@ -742,6 +781,102 @@ def _register_forward_hook_or_replace_fn(
raise ValueError(f"Unknown mode {mode!r}") 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 ------------------------------------------ # -------------------------------------- util fn ------------------------------------------
@@ -1172,6 +1307,61 @@ def _get_local_ip_by_remote() -> Optional[str]:
return None 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 ------------------------------------------ # -------------------------------------- framework plugins ------------------------------------------
+310
View File
@@ -4,8 +4,10 @@ import os
import sys import sys
import threading import threading
import time import time
import traceback
from contextlib import contextmanager from contextlib import contextmanager
from pathlib import Path from pathlib import Path
from typing import Optional
import pytest import pytest
import requests import requests
@@ -20,6 +22,7 @@ from sglang.srt.debug_utils.dumper import (
_Dumper, _Dumper,
_format_tags, _format_tags,
_get_default_exp_name, _get_default_exp_name,
_Grafter,
_log, _log,
_map_tensor, _map_tensor,
_materialize_value, _materialize_value,
@@ -2592,5 +2595,312 @@ class TestRecomputeStatus:
assert _detect_recompute_status() == _RecomputeStatus.DISABLED 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__": if __name__ == "__main__":
sys.exit(pytest.main([__file__])) sys.exit(pytest.main([__file__]))