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"
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 ------------------------------------------