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