Support multi-rank exchange via all_gather_object in dumper grafter (#24510)

This commit is contained in:
fzyzcjy
2026-05-06 16:57:20 +08:00
committed by GitHub
parent ebd64f5d40
commit 833279eb2e
2 changed files with 363 additions and 22 deletions
+60 -22
View File
@@ -842,30 +842,36 @@ class _Grafter:
role = _GraftRole(cfg.grafter_role)
is_send = self._is_sender(role=role, direction=direction)
# 1+1 broadcast: sender side ships the tensor as a pickled object;
# recv side feeds it through the user transform (default: identity)
# and `value.copy_()` the result.
sender_rank = 0 if direction == _GraftDirection.B2T else 1
obj_list: list = [None]
# all-gather over the graft world; sender ranks contribute `value`,
# recv ranks contribute None (their local target is private and
# shouldn't leak). all_gather_object is pickle-routed, so tensor
# shapes may differ across sender ranks.
total_world = cfg.grafter_baseline_world_size + cfg.grafter_target_world_size
my_contribution = value if is_send else None
gathered: list = [None] * total_world
dist.all_gather_object(gathered, my_contribution, group=self._pg)
if is_send:
obj_list = [value]
_log(f"[Grafter] send role={role.value} dir={direction.value} tags={tags}")
dist.broadcast_object_list(obj_list, src=sender_rank, group=self._pg)
if is_send:
return
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)
sender_contribs = self._sender_slice(direction=direction, gathered=gathered)
# Pickled CUDA tensors restore to their original-device name; that
# may not match this process's local device, so normalize.
sender_tensors = [
(t.to(value.device) if isinstance(t, torch.Tensor) else t)
for t in sender_contribs
]
# Transform + copy_ are wrapped: a buggy user transform must NOT
# crash the whole training/inference run. On error we log the full
# traceback and skip this graft point; downstream sees the recv
# side's original tensor unchanged.
try:
value_to_override = self._apply_transform([received], target=value)
_log(f"[Grafter] recv role={role.value} dir={direction.value} tags={tags}")
value_to_override = self._apply_transform(sender_tensors, target=value)
_log(
f"[Grafter] recv role={role.value} dir={direction.value} "
f"tags={tags} n_senders={len(sender_tensors)}"
)
value.copy_(value_to_override)
except Exception as e:
_log(
@@ -875,6 +881,12 @@ class _Grafter:
f"{traceback.format_exc()}"
)
def _sender_slice(self, *, direction: "_GraftDirection", gathered: list) -> list:
cfg = self._config
if direction == _GraftDirection.B2T:
return gathered[: cfg.grafter_baseline_world_size]
return gathered[cfg.grafter_baseline_world_size :]
def _apply_transform(
self,
received_list: list,
@@ -890,16 +902,27 @@ class _Grafter:
def _default_transform(
received_list: list, *, target: torch.Tensor
) -> torch.Tensor:
"""Identity-by-rank fallback. For the 1+1 setup currently supported,
just returns the single received tensor; requires shape match."""
candidate = received_list[0]
"""Identity-by-rank fallback. Requires #senders == #recvs and
shape(received_list[my_recv_rank]) == shape(target). Otherwise raises
and asks the user for a transform."""
my_recv_rank = dist.get_rank()
recv_world_size = dist.get_world_size()
if len(received_list) != recv_world_size:
raise RuntimeError(
f"[Grafter] no grafter_transform_path set; default "
f"identity-by-rank requires #senders == #recvs but got "
f"#senders={len(received_list)} vs #recvs={recv_world_size}. "
f"Provide a transform via "
f"DUMPER_GRAFTER_TRANSFORM_PATH=pkg.module.symbol."
)
candidate = received_list[my_recv_rank]
if candidate.shape != target.shape:
raise RuntimeError(
f"[Grafter] no grafter_transform_path set; default "
f"identity-by-rank requires matching shapes but "
f"received_list[0].shape={tuple(candidate.shape)} != "
f"target.shape={tuple(target.shape)}. Provide a transform via "
f"DUMPER_GRAFTER_TRANSFORM_PATH=pkg.module.symbol."
f"received_list[{my_recv_rank}].shape={tuple(candidate.shape)} "
f"!= target.shape={tuple(target.shape)}. Provide a transform "
f"via DUMPER_GRAFTER_TRANSFORM_PATH=pkg.module.symbol."
)
return candidate
@@ -938,11 +961,26 @@ class _Grafter:
dist.is_initialized()
), "[Grafter] default torch.distributed must be initialized"
role = _GraftRole(cfg.grafter_role)
global_rank = 0 if role == _GraftRole.BASELINE else 1
local_world = dist.get_world_size()
local_rank = dist.get_rank()
if role == _GraftRole.BASELINE:
assert local_world == cfg.grafter_baseline_world_size, (
f"[Grafter] grafter_baseline_world_size={cfg.grafter_baseline_world_size} "
f"but dist.get_world_size()={local_world}; they must match on the baseline side"
)
global_rank = local_rank
else:
assert local_world == cfg.grafter_target_world_size, (
f"[Grafter] grafter_target_world_size={cfg.grafter_target_world_size} "
f"but dist.get_world_size()={local_world}; they must match on the target side"
)
global_rank = cfg.grafter_baseline_world_size + local_rank
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"baseline_world={cfg.grafter_baseline_world_size} "
f"target_world={cfg.grafter_target_world_size} "
f"rank={global_rank} init_method={init_method} "
f"backend={cfg.grafter_backend} name={cfg.grafter_group_name}"
)
+303
View File
@@ -3246,5 +3246,308 @@ class TestGrafterDistributed:
dist.destroy_process_group(grafter._pg)
def _run_graft_test_cpu_multi(
worker_func, *, baseline_world: int, target_world: int, **kwargs
):
"""Spawn (baseline_world + target_world) CPU-only processes (gloo backend).
Used to exercise asymmetric multi-rank cases (e.g. 4 baseline ranks and
2 target ranks) that we can't run on the 2-GPU CI fleet. Each role gets
its OWN default PG (gloo, world=role_world); the graft cross-system PG
spans all ranks.
The worker function receives (role, local_rank, **kwargs).
"""
import torch.multiprocessing as mp
role_ports = {
"baseline": find_available_port(29800),
"target": find_available_port(29900),
}
ctx = mp.get_context("spawn")
result_queue = ctx.Queue()
processes = []
total = baseline_world + target_world
for global_rank in range(total):
if global_rank < baseline_world:
role = "baseline"
local_rank = global_rank
local_world = baseline_world
else:
role = "target"
local_rank = global_rank - baseline_world
local_world = target_world
p = ctx.Process(
target=_graft_cpu_worker_entry,
args=(
role,
local_rank,
local_world,
role_ports[role],
worker_func,
result_queue,
kwargs,
),
)
p.start()
processes.append(p)
for p in processes:
p.join()
errors = [result_queue.get() for _ in range(total)]
errors = [e for e in errors if e]
if errors:
raise AssertionError("\n".join(errors))
def _graft_cpu_worker_entry(
role, local_rank, local_world, port, worker_func, result_queue, kwargs
):
dist.init_process_group(
backend="gloo",
init_method=f"tcp://127.0.0.1:{port}",
world_size=local_world,
rank=local_rank,
)
try:
worker_func(role=role, local_rank=local_rank, **kwargs)
result_queue.put(None)
except Exception as e:
result_queue.put(
f"role={role} local_rank={local_rank}: {e}\n{traceback.format_exc()}"
)
finally:
dist.destroy_process_group()
def _make_multi_rank_config(
*,
role: str,
graft_port: int,
group_name: str,
baseline_world: int,
target_world: int,
transform_path: Optional[str],
direction: str,
) -> DumperConfig:
return DumperConfig(
grafter_enable=True,
grafter_role=role,
grafter_b2t_filter="name == 'x'" if direction == "b2t" else None,
grafter_t2b_filter="name == 'x'" if direction == "t2b" else None,
grafter_master_address="127.0.0.1",
grafter_master_port=graft_port,
grafter_baseline_world_size=baseline_world,
grafter_target_world_size=target_world,
grafter_backend="gloo",
grafter_group_name=group_name,
grafter_timeout=30,
grafter_transform_path=transform_path,
)
class TestGrafterMultiRankCpu:
"""Coverage of asymmetric multi-rank cases via CPU/gloo (CI fleet has
only 2 GPUs, which is too few for these cases)."""
def test_4_baseline_2_target_b2t_with_user_transform(self, tmp_path: Path):
"""4 baseline senders -> 2 target receivers via b2t graft."""
module_name = "_xform_assert_4_senders"
(tmp_path / f"{module_name}.py").write_text(
"import torch\n"
"def transform(received_list, target):\n"
" rl = received_list\n"
" assert len(rl) == 4, f'expected 4 senders, got {len(rl)}'\n"
" for i, t in enumerate(rl):\n"
" v = float(t.flatten()[0].item())\n"
" assert v == float(i), f'rl[{i}][0]={v}, want {float(i)}'\n"
" return torch.full_like(target, 999.0)\n"
)
graft_port = find_available_port(29655)
_run_graft_test_cpu_multi(
self._test_4b_2t_func,
baseline_world=4,
target_world=2,
graft_port=graft_port,
group_name="grafter_4b_2t",
transform_dir=str(tmp_path),
transform_path=f"{module_name}.transform",
)
@staticmethod
def _test_4b_2t_func(
role, local_rank, graft_port, group_name, transform_dir, transform_path
):
sys.path.insert(0, transform_dir)
cfg = _make_multi_rank_config(
role=role,
graft_port=graft_port,
group_name=group_name,
baseline_world=4,
target_world=2,
transform_path=transform_path,
direction="b2t",
)
grafter = _Grafter(config=cfg)
try:
if role == "baseline":
tensor = torch.full((3,), float(local_rank))
grafter.maybe_intercept(value=tensor, tags={"name": "x"})
else:
target = torch.full((3,), 99.0)
grafter.maybe_intercept(value=target, tags={"name": "x"})
assert target.tolist() == [999.0, 999.0, 999.0], target.tolist()
finally:
if grafter._pg is not None:
dist.destroy_process_group(grafter._pg)
def test_2_target_4_baseline_t2b_with_user_transform(self, tmp_path: Path):
"""Mirror image: 2 target senders -> 4 baseline receivers."""
module_name = "_xform_assert_2_senders_t2b"
(tmp_path / f"{module_name}.py").write_text(
"import torch\n"
"def transform(received_list, target):\n"
" rl = received_list\n"
" assert len(rl) == 2, f'expected 2 senders, got {len(rl)}'\n"
" for i, t in enumerate(rl):\n"
" v = float(t.flatten()[0].item())\n"
" assert v == float(i + 100), f'rl[{i}][0]={v}'\n"
" return torch.full_like(target, 7.0)\n"
)
graft_port = find_available_port(29670)
_run_graft_test_cpu_multi(
self._test_2t_4b_func,
baseline_world=4,
target_world=2,
graft_port=graft_port,
group_name="grafter_2t_4b",
transform_dir=str(tmp_path),
transform_path=f"{module_name}.transform",
)
@staticmethod
def _test_2t_4b_func(
role, local_rank, graft_port, group_name, transform_dir, transform_path
):
sys.path.insert(0, transform_dir)
cfg = _make_multi_rank_config(
role=role,
graft_port=graft_port,
group_name=group_name,
baseline_world=4,
target_world=2,
transform_path=transform_path,
direction="t2b",
)
grafter = _Grafter(config=cfg)
try:
if role == "target":
tensor = torch.full((3,), float(local_rank + 100))
grafter.maybe_intercept(value=tensor, tags={"name": "x"})
else:
target = torch.full((3,), 99.0)
grafter.maybe_intercept(value=target, tags={"name": "x"})
assert target.tolist() == [7.0, 7.0, 7.0], target.tolist()
finally:
if grafter._pg is not None:
dist.destroy_process_group(grafter._pg)
def test_default_transform_with_asymmetric_world_logs_and_skips(self):
"""Default identity-by-rank requires #senders == #recvs; with 4 vs 2
and no user transform, recv catches RuntimeError + leaves target
unchanged."""
graft_port = find_available_port(29675)
_run_graft_test_cpu_multi(
self._test_default_asym_func,
baseline_world=4,
target_world=2,
graft_port=graft_port,
group_name="grafter_default_asym",
)
@staticmethod
def _test_default_asym_func(role, local_rank, graft_port, group_name):
cfg = _make_multi_rank_config(
role=role,
graft_port=graft_port,
group_name=group_name,
baseline_world=4,
target_world=2,
transform_path=None,
direction="b2t",
)
grafter = _Grafter(config=cfg)
try:
if role == "baseline":
tensor = torch.full((3,), float(local_rank))
grafter.maybe_intercept(value=tensor, tags={"name": "x"})
else:
target = torch.full((3,), 42.0)
with _capture_stdout() as captured:
grafter.maybe_intercept(value=target, tags={"name": "x"})
assert target.tolist() == [42.0, 42.0, 42.0], target.tolist()
output = captured.getvalue()
assert "transform/copy_ raised RuntimeError" in output, output
assert "#senders=4" in output and "#recvs=2" in output, output
assert "Traceback (most recent call last)" in output, output
finally:
if grafter._pg is not None:
dist.destroy_process_group(grafter._pg)
def test_mixed_shape_senders_via_user_transform(self, tmp_path: Path):
"""`all_gather_object` is pickle-routed, so sender ranks may
contribute tensors with DIFFERENT shapes. The user transform handles
the dispatch."""
module_name = "_xform_concat_mixed_shape"
(tmp_path / f"{module_name}.py").write_text(
"import torch\n"
"def transform(received_list, target):\n"
" expected_shapes = [(i + 1,) for i in range(len(received_list))]\n"
" actual_shapes = [tuple(t.shape) for t in received_list]\n"
" assert actual_shapes == expected_shapes, actual_shapes\n"
" return torch.cat(received_list)\n"
)
graft_port = find_available_port(29680)
_run_graft_test_cpu_multi(
self._test_mixed_shape_func,
baseline_world=4,
target_world=2,
graft_port=graft_port,
group_name="grafter_mixed_shape",
transform_dir=str(tmp_path),
transform_path=f"{module_name}.transform",
)
@staticmethod
def _test_mixed_shape_func(
role, local_rank, graft_port, group_name, transform_dir, transform_path
):
sys.path.insert(0, transform_dir)
cfg = _make_multi_rank_config(
role=role,
graft_port=graft_port,
group_name=group_name,
baseline_world=4,
target_world=2,
transform_path=transform_path,
direction="b2t",
)
grafter = _Grafter(config=cfg)
try:
if role == "baseline":
tensor = torch.full((local_rank + 1,), float(local_rank))
grafter.maybe_intercept(value=tensor, tags={"name": "x"})
else:
target = torch.zeros(10)
grafter.maybe_intercept(value=target, tags={"name": "x"})
expected = [0.0] + [1.0] * 2 + [2.0] * 3 + [3.0] * 4
assert target.tolist() == expected, target.tolist()
finally:
if grafter._pg is not None:
dist.destroy_process_group(grafter._pg)
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))