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_backend: str = "nccl"
|
||||||
grafter_group_name: str = "graft"
|
grafter_group_name: str = "graft"
|
||||||
grafter_timeout: int = 300
|
grafter_timeout: int = 300
|
||||||
# Fully-qualified Python path "pkg.subpkg.module.fn_name". When set, the
|
# Fully-qualified Python path "pkg.subpkg.module.fn_name"
|
||||||
# recv side calls this function with (received_list, target) and copies
|
# None -> use the default identity-by-rank fallback in _Grafter._default_transform.
|
||||||
# the result into target. None -> use the default identity-by-rank
|
|
||||||
# fallback in `_Grafter._default_transform`.
|
|
||||||
grafter_transform_path: Optional[str] = None
|
grafter_transform_path: Optional[str] = None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -831,27 +829,34 @@ class GraftTransformInput:
|
|||||||
|
|
||||||
|
|
||||||
class _Grafter:
|
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 ->
|
Both sides set the SAME grafter_b2t_filter (names that flow baseline ->
|
||||||
target) and `grafter_t2b_filter` (names that flow target -> baseline).
|
target) and grafter_t2b_filter (names that flow target -> baseline). The
|
||||||
The only per-side difference is `grafter_role`, which tells the side
|
only per-side difference is grafter_role ("baseline" | "target"), which
|
||||||
whether it's the sender or the receiver for the matched direction.
|
determines whether a name match means send or recv on this side.
|
||||||
Receiver overwrites its local target tensor with the sender's via
|
|
||||||
`value.copy_()`.
|
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._config = config
|
||||||
self._pg: Optional[dist.ProcessGroup] = None
|
self._pg = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def enabled(self) -> bool:
|
||||||
|
return self._config.grafter_enable
|
||||||
|
|
||||||
def maybe_intercept(
|
def maybe_intercept(
|
||||||
self,
|
self, *, value: Any, tags: dict, extras: Optional[dict] = None
|
||||||
*,
|
|
||||||
value,
|
|
||||||
tags: dict,
|
|
||||||
extras: Optional[dict] = None,
|
|
||||||
) -> 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
|
cfg = self._config
|
||||||
if not cfg.grafter_enable:
|
if not cfg.grafter_enable:
|
||||||
return
|
return
|
||||||
@@ -864,9 +869,9 @@ class _Grafter:
|
|||||||
_log(
|
_log(
|
||||||
f"[Grafter] tags={tags} matched grafter_{direction.value}_filter but "
|
f"[Grafter] tags={tags} matched grafter_{direction.value}_filter but "
|
||||||
f"value is not a torch.Tensor (got type={type(value).__name__}); "
|
f"value is not a torch.Tensor (got type={type(value).__name__}); "
|
||||||
f"skipping graft. Common cause: dumper.dump called with a "
|
f"skipping graft. Common cause: dumper.dump called with a non-tensor "
|
||||||
f"non-tensor value (dict, list, ...) on this name. Either "
|
f"value (dict, list, ...) on this name. Either narrow the filter or "
|
||||||
f"narrow the filter or wrap the value in a tensor."
|
f"wrap the value in a tensor."
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -876,7 +881,8 @@ class _Grafter:
|
|||||||
|
|
||||||
# all-gather over the graft world; sender ranks contribute (value,
|
# all-gather over the graft world; sender ranks contribute (value,
|
||||||
# extras) tuples, recv ranks contribute None (their local target is
|
# 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
|
total_world = cfg.grafter_baseline_world_size + cfg.grafter_target_world_size
|
||||||
my_contribution = (value, extras) if is_send else None
|
my_contribution = (value, extras) if is_send else None
|
||||||
gathered: list = [None] * total_world
|
gathered: list = [None] * total_world
|
||||||
@@ -890,8 +896,8 @@ class _Grafter:
|
|||||||
return
|
return
|
||||||
|
|
||||||
sender_contribs = self._sender_slice(direction=direction, gathered=gathered)
|
sender_contribs = self._sender_slice(direction=direction, gathered=gathered)
|
||||||
# Pickled CUDA tensors restore to their original-device name; that
|
# Pickled CUDA tensors are restored on their original-device name;
|
||||||
# may not match this process's local device, so normalize.
|
# that may not match this process's local device, so normalize.
|
||||||
sender_tensors = [
|
sender_tensors = [
|
||||||
(c[0].to(value.device) if isinstance(c[0], torch.Tensor) else c[0])
|
(c[0].to(value.device) if isinstance(c[0], torch.Tensor) else c[0])
|
||||||
for c in sender_contribs
|
for c in sender_contribs
|
||||||
@@ -928,66 +934,13 @@ class _Grafter:
|
|||||||
f"{traceback.format_exc()}"
|
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"]:
|
def _classify_direction(self, tags: dict) -> Optional["_GraftDirection"]:
|
||||||
cfg = self._config
|
cfg = self._config
|
||||||
match_b2t = self._match(cfg.grafter_b2t_filter, tags)
|
match_b2t = self._match(cfg.grafter_b2t_filter, tags)
|
||||||
match_t2b = self._match(cfg.grafter_t2b_filter, tags)
|
match_t2b = self._match(cfg.grafter_t2b_filter, tags)
|
||||||
if match_b2t and match_t2b:
|
if match_b2t and match_t2b:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"[Grafter] tags={tags} matched BOTH grafter_b2t_filter "
|
f"[Grafter] tags={tags} matched BOTH grafter_b2t_filter and grafter_t2b_filter"
|
||||||
f"and grafter_t2b_filter"
|
|
||||||
)
|
)
|
||||||
if match_b2t:
|
if match_b2t:
|
||||||
return _GraftDirection.B2T
|
return _GraftDirection.B2T
|
||||||
@@ -1000,6 +953,12 @@ class _Grafter:
|
|||||||
# baseline is the sender for B2T names; target is the sender for T2B.
|
# baseline is the sender for B2T names; target is the sender for T2B.
|
||||||
return (role == _GraftRole.BASELINE) == (direction == _GraftDirection.B2T)
|
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
|
@staticmethod
|
||||||
def _match(expr: Optional[str], tags: dict) -> bool:
|
def _match(expr: Optional[str], tags: dict) -> bool:
|
||||||
if expr is None:
|
if expr is None:
|
||||||
@@ -1050,6 +1009,62 @@ class _Grafter:
|
|||||||
timeout_seconds=cfg.grafter_timeout,
|
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 ------------------------------------------
|
# -------------------------------------- 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;
|
sglang.srt.debug_utils.dump_comparator._compute_and_print_diff;
|
||||||
intentionally inlined here to keep dumper.py free of cross-file imports.
|
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."""
|
enough for the order-of-magnitude diff summary we log."""
|
||||||
if a.shape != b.shape:
|
if a.shape != b.shape:
|
||||||
return f"shape mismatch (a={tuple(a.shape)} vs b={tuple(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
|
return None
|
||||||
|
|
||||||
|
|
||||||
@functools.lru_cache(maxsize=None)
|
|
||||||
def _load_function(path: str) -> Callable:
|
def _load_function(path: str) -> Callable:
|
||||||
"""Resolve a fully-qualified Python path 'pkg.module.symbol' to its object.
|
"""Resolve a fully-qualified Python path 'pkg.module.symbol' to its object.
|
||||||
|
|
||||||
Copied (verbatim, minus the function-registry branch) from
|
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.
|
cross-package dependency.
|
||||||
"""
|
"""
|
||||||
import importlib
|
import importlib
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user