Support per-call extras and dataclass transform input in dumper grafter (#24511)
This commit is contained in:
@@ -279,6 +279,7 @@ class _Dumper:
|
||||
save: bool = True,
|
||||
dims: Optional[str] = None,
|
||||
dims_grad: Optional[str] = None,
|
||||
grafter_extras: Optional[dict] = None,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
value_meta: dict = {}
|
||||
@@ -302,6 +303,7 @@ class _Dumper:
|
||||
grad_tag="Dumper.Grad",
|
||||
value_meta_only_fields=value_meta,
|
||||
grad_meta_only_fields=grad_meta,
|
||||
grafter_extras=grafter_extras,
|
||||
)
|
||||
|
||||
def dump_model(
|
||||
@@ -458,6 +460,7 @@ class _Dumper:
|
||||
grad_tag: str,
|
||||
value_meta_only_fields: Optional[dict] = None,
|
||||
grad_meta_only_fields: Optional[dict] = None,
|
||||
grafter_extras: Optional[dict] = None,
|
||||
) -> None:
|
||||
self._http_manager # noqa: B018
|
||||
|
||||
@@ -480,7 +483,7 @@ class _Dumper:
|
||||
|
||||
recompute_meta = recompute_status.to_pseudo_parallel_meta()
|
||||
value = _materialize_value(value)
|
||||
self._grafter.maybe_intercept(value=value, tags=tags)
|
||||
self._grafter.maybe_intercept(value=value, tags=tags, extras=grafter_extras)
|
||||
|
||||
if enable_value:
|
||||
self._dump_single(
|
||||
@@ -804,6 +807,29 @@ class _GraftDirection(enum.Enum):
|
||||
T2B = "t2b" # name flows target -> baseline
|
||||
|
||||
|
||||
@dataclass
|
||||
class GraftTransformInput:
|
||||
"""Single argument passed to a user-supplied transform function.
|
||||
|
||||
User transforms have signature::
|
||||
|
||||
def transform(graft_input: GraftTransformInput) -> torch.Tensor: ...
|
||||
|
||||
The dataclass shape lets us add fields (e.g., direction, sender ranks)
|
||||
later without breaking existing transforms.
|
||||
"""
|
||||
|
||||
# Full dumper.dump tags dict (name + recompute_status + extra_kwargs + ctx).
|
||||
tags: "dict[str, Any]"
|
||||
# One tensor per sender rank, in sender-rank order.
|
||||
received_list: "list[torch.Tensor]"
|
||||
# Parallel list of per-sender `grafter_extras` (the dict passed to
|
||||
# dumper.dump on each sender; None if the sender omitted it).
|
||||
received_extras_list: "list[Optional[dict]]"
|
||||
# Recv side's local tensor that will be copy_'d into.
|
||||
target: "torch.Tensor"
|
||||
|
||||
|
||||
class _Grafter:
|
||||
"""1+1 cross-system tensor grafter.
|
||||
|
||||
@@ -819,7 +845,13 @@ class _Grafter:
|
||||
self._config = config
|
||||
self._pg: Optional[dist.ProcessGroup] = None
|
||||
|
||||
def maybe_intercept(self, *, value, tags: dict) -> None:
|
||||
def maybe_intercept(
|
||||
self,
|
||||
*,
|
||||
value,
|
||||
tags: dict,
|
||||
extras: Optional[dict] = None,
|
||||
) -> None:
|
||||
cfg = self._config
|
||||
if not cfg.grafter_enable:
|
||||
return
|
||||
@@ -842,35 +874,45 @@ class _Grafter:
|
||||
role = _GraftRole(cfg.grafter_role)
|
||||
is_send = self._is_sender(role=role, direction=direction)
|
||||
|
||||
# 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.
|
||||
# all-gather over the graft world; sender ranks contribute (value,
|
||||
# extras) tuples, recv ranks contribute None (their local target is
|
||||
# private and shouldn't leak).
|
||||
total_world = cfg.grafter_baseline_world_size + cfg.grafter_target_world_size
|
||||
my_contribution = value if is_send else None
|
||||
my_contribution = (value, extras) if is_send else None
|
||||
gathered: list = [None] * total_world
|
||||
dist.all_gather_object(gathered, my_contribution, group=self._pg)
|
||||
|
||||
if is_send:
|
||||
_log(f"[Grafter] send role={role.value} dir={direction.value} tags={tags}")
|
||||
_log(
|
||||
f"[Grafter] send role={role.value} dir={direction.value} "
|
||||
f"tags={tags} extras={extras}"
|
||||
)
|
||||
return
|
||||
|
||||
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
|
||||
(c[0].to(value.device) if isinstance(c[0], torch.Tensor) else c[0])
|
||||
for c in sender_contribs
|
||||
]
|
||||
sender_extras = [c[1] for c 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(sender_tensors, target=value)
|
||||
value_to_override = self._apply_transform(
|
||||
tags=tags,
|
||||
received_list=sender_tensors,
|
||||
received_extras_list=sender_extras,
|
||||
target=value,
|
||||
)
|
||||
_log(
|
||||
f"[Grafter] recv role={role.value} dir={direction.value} "
|
||||
f"tags={tags} n_senders={len(sender_tensors)}"
|
||||
f"tags={tags} n_senders={len(sender_tensors)} "
|
||||
f"sender_extras={sender_extras}"
|
||||
)
|
||||
value.copy_(value_to_override)
|
||||
except Exception as e:
|
||||
@@ -889,22 +931,29 @@ class _Grafter:
|
||||
|
||||
def _apply_transform(
|
||||
self,
|
||||
received_list: list,
|
||||
*,
|
||||
tags: dict,
|
||||
received_list: list,
|
||||
received_extras_list: list,
|
||||
target: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
graft_input = GraftTransformInput(
|
||||
tags=tags,
|
||||
received_list=received_list,
|
||||
received_extras_list=received_extras_list,
|
||||
target=target,
|
||||
)
|
||||
path = self._config.grafter_transform_path
|
||||
if path is None:
|
||||
return self._default_transform(received_list, target=target)
|
||||
return _load_function(path)(received_list, target)
|
||||
fn = self._default_transform if path is None else _load_function(path)
|
||||
return fn(graft_input)
|
||||
|
||||
@staticmethod
|
||||
def _default_transform(
|
||||
received_list: list, *, target: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
def _default_transform(graft_input: GraftTransformInput) -> torch.Tensor:
|
||||
"""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."""
|
||||
received_list = graft_input.received_list
|
||||
target = graft_input.target
|
||||
my_recv_rank = dist.get_rank()
|
||||
recv_world_size = dist.get_world_size()
|
||||
if len(received_list) != recv_world_size:
|
||||
|
||||
Reference in New Issue
Block a user