Add e2e test with log snapshot in dumper grafter (#24513)
This commit is contained in:
@@ -152,10 +152,8 @@ class DumperConfig(_BaseConfig):
|
||||
grafter_backend: str = "nccl"
|
||||
grafter_group_name: str = "graft"
|
||||
grafter_timeout: int = 300
|
||||
# Fully-qualified Python path "pkg.subpkg.module.fn_name". When set, the
|
||||
# recv side calls this function with (received_list, target) and copies
|
||||
# the result into target. None -> use the default identity-by-rank
|
||||
# fallback in `_Grafter._default_transform`.
|
||||
# Fully-qualified Python path "pkg.subpkg.module.fn_name"
|
||||
# None -> use the default identity-by-rank fallback in _Grafter._default_transform.
|
||||
grafter_transform_path: Optional[str] = None
|
||||
|
||||
@classmethod
|
||||
@@ -831,27 +829,34 @@ class GraftTransformInput:
|
||||
|
||||
|
||||
class _Grafter:
|
||||
"""1+1 cross-system tensor grafter.
|
||||
"""Cross-system tensor transplant. Triggered silently from dumper.dump.
|
||||
|
||||
Both sides set the SAME `grafter_b2t_filter` (names that flow baseline ->
|
||||
target) and `grafter_t2b_filter` (names that flow target -> baseline).
|
||||
The only per-side difference is `grafter_role`, which tells the side
|
||||
whether it's the sender or the receiver for the matched direction.
|
||||
Receiver overwrites its local target tensor with the sender's via
|
||||
`value.copy_()`.
|
||||
Both sides set the SAME grafter_b2t_filter (names that flow baseline ->
|
||||
target) and grafter_t2b_filter (names that flow target -> baseline). The
|
||||
only per-side difference is grafter_role ("baseline" | "target"), which
|
||||
determines whether a name match means send or recv on this side.
|
||||
|
||||
Graft global rank layout: baseline occupies ranks 0..baseline_world-1;
|
||||
target occupies ranks baseline_world..baseline_world+target_world-1. Each
|
||||
side derives its own rank from its local default PG via dist.get_rank().
|
||||
|
||||
Please refer to TestGrafterE2eExample in tests for an example.
|
||||
"""
|
||||
|
||||
def __init__(self, *, config: DumperConfig) -> None:
|
||||
def __init__(self, *, config: DumperConfig):
|
||||
self._config = config
|
||||
self._pg: Optional[dist.ProcessGroup] = None
|
||||
self._pg = None
|
||||
|
||||
@property
|
||||
def enabled(self) -> bool:
|
||||
return self._config.grafter_enable
|
||||
|
||||
def maybe_intercept(
|
||||
self,
|
||||
*,
|
||||
value,
|
||||
tags: dict,
|
||||
extras: Optional[dict] = None,
|
||||
self, *, value: Any, tags: dict, extras: Optional[dict] = None
|
||||
) -> None:
|
||||
"""Intercept a dumper.dump call. `extras` is per-call auxiliary data
|
||||
(e.g., shard layout, dtype hint) that the sender attaches and the
|
||||
recv side's transform receives as `received_extras_list`."""
|
||||
cfg = self._config
|
||||
if not cfg.grafter_enable:
|
||||
return
|
||||
@@ -864,9 +869,9 @@ class _Grafter:
|
||||
_log(
|
||||
f"[Grafter] tags={tags} matched grafter_{direction.value}_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."
|
||||
f"skipping graft. Common cause: dumper.dump called with a non-tensor "
|
||||
f"value (dict, list, ...) on this name. Either narrow the filter or "
|
||||
f"wrap the value in a tensor."
|
||||
)
|
||||
return
|
||||
|
||||
@@ -876,7 +881,8 @@ class _Grafter:
|
||||
|
||||
# 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).
|
||||
# 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, extras) if is_send else None
|
||||
gathered: list = [None] * total_world
|
||||
@@ -890,8 +896,8 @@ class _Grafter:
|
||||
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.
|
||||
# Pickled CUDA tensors are restored on their original-device name;
|
||||
# that may not match this process's local device, so normalize.
|
||||
sender_tensors = [
|
||||
(c[0].to(value.device) if isinstance(c[0], torch.Tensor) else c[0])
|
||||
for c in sender_contribs
|
||||
@@ -928,66 +934,13 @@ 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,
|
||||
*,
|
||||
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
|
||||
fn = self._default_transform if path is None else _load_function(path)
|
||||
return fn(graft_input)
|
||||
|
||||
@staticmethod
|
||||
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:
|
||||
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[{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
|
||||
|
||||
def _classify_direction(self, tags: dict) -> Optional["_GraftDirection"]:
|
||||
cfg = self._config
|
||||
match_b2t = self._match(cfg.grafter_b2t_filter, tags)
|
||||
match_t2b = self._match(cfg.grafter_t2b_filter, tags)
|
||||
if match_b2t and match_t2b:
|
||||
raise RuntimeError(
|
||||
f"[Grafter] tags={tags} matched BOTH grafter_b2t_filter "
|
||||
f"and grafter_t2b_filter"
|
||||
f"[Grafter] tags={tags} matched BOTH grafter_b2t_filter and grafter_t2b_filter"
|
||||
)
|
||||
if match_b2t:
|
||||
return _GraftDirection.B2T
|
||||
@@ -1000,6 +953,12 @@ class _Grafter:
|
||||
# baseline is the sender for B2T names; target is the sender for T2B.
|
||||
return (role == _GraftRole.BASELINE) == (direction == _GraftDirection.B2T)
|
||||
|
||||
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 :]
|
||||
|
||||
@staticmethod
|
||||
def _match(expr: Optional[str], tags: dict) -> bool:
|
||||
if expr is None:
|
||||
@@ -1050,6 +1009,62 @@ class _Grafter:
|
||||
timeout_seconds=cfg.grafter_timeout,
|
||||
)
|
||||
|
||||
def _apply_transform(
|
||||
self,
|
||||
*,
|
||||
tags: dict,
|
||||
received_list: list,
|
||||
received_extras_list: list,
|
||||
target: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
# TODO: integrate with dump_comparator unsharder annotations once
|
||||
# full inverse (sharded -> global -> sharded) transforms exist.
|
||||
graft_input = GraftTransformInput(
|
||||
tags=tags,
|
||||
received_list=received_list,
|
||||
received_extras_list=received_extras_list,
|
||||
target=target,
|
||||
)
|
||||
path = self._config.grafter_transform_path
|
||||
fn = self._default_transform if path is None else _load_function(path)
|
||||
return fn(graft_input)
|
||||
|
||||
@staticmethod
|
||||
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:
|
||||
raise RuntimeError(
|
||||
_Grafter._default_transform_error(
|
||||
f"requires #senders == #recvs but got "
|
||||
f"#senders={len(received_list)} vs #recvs={recv_world_size}"
|
||||
)
|
||||
)
|
||||
candidate = received_list[my_recv_rank]
|
||||
if candidate.shape != target.shape:
|
||||
raise RuntimeError(
|
||||
_Grafter._default_transform_error(
|
||||
f"requires matching shapes but "
|
||||
f"received_list[{my_recv_rank}].shape={tuple(candidate.shape)} "
|
||||
f"!= target.shape={tuple(target.shape)}"
|
||||
)
|
||||
)
|
||||
return candidate
|
||||
|
||||
@staticmethod
|
||||
def _default_transform_error(detail: str) -> str:
|
||||
return (
|
||||
f"[Grafter] no grafter_transform_path set; default identity-by-rank "
|
||||
f"{detail}. Provide a transform via "
|
||||
f"DUMPER_GRAFTER_TRANSFORM_PATH=pkg.module.symbol defining "
|
||||
f"`transform(graft_input: GraftTransformInput) -> Tensor`."
|
||||
)
|
||||
|
||||
|
||||
# -------------------------------------- util fn ------------------------------------------
|
||||
|
||||
@@ -1186,7 +1201,7 @@ def _compare_tensors_quick(a: "torch.Tensor", b: "torch.Tensor") -> str:
|
||||
sglang.srt.debug_utils.dump_comparator._compute_and_print_diff;
|
||||
intentionally inlined here to keep dumper.py free of cross-file imports.
|
||||
|
||||
Different dtypes are fine -- we unify by casting both to fp32, which is
|
||||
Different dtypes are fine — we unify by casting both to fp32, which is
|
||||
enough for the order-of-magnitude diff summary we log."""
|
||||
if a.shape != b.shape:
|
||||
return f"shape mismatch (a={tuple(a.shape)} vs b={tuple(b.shape)})"
|
||||
@@ -1510,12 +1525,11 @@ def _get_local_ip_by_remote() -> Optional[str]:
|
||||
return None
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=None)
|
||||
def _load_function(path: str) -> Callable:
|
||||
"""Resolve a fully-qualified Python path 'pkg.module.symbol' to its object.
|
||||
|
||||
Copied (verbatim, minus the function-registry branch) from
|
||||
miles.utils.misc.load_function -- kept inline so dumper.py has no
|
||||
miles.utils.misc.load_function — kept inline so dumper.py has no
|
||||
cross-package dependency.
|
||||
"""
|
||||
import importlib
|
||||
|
||||
Reference in New Issue
Block a user