Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
20 commits
Select commit Hold shift + click to select a range
bc91469
feat(rollout): generic collective_rpc across all DP and TP ranks
Sep 18, 2026
776d740
fix(rollout): answer every sync caller, and never a fire-and-forget one
Sep 28, 2026
a584997
fix(rollout): keep collective_rpc clear of main's block-table codec
Sep 28, 2026
f9a7aa1
fix(rollout): refuse collective_rpc when engines are not DP ranks
Sep 28, 2026
7088cfa
fix(rollout): refuse a non-string collective_rpc method name
Sep 28, 2026
8626ef1
fix(rollout): release a barrier call's survivors when a rank dies
Sep 29, 2026
1b2cbbb
fix(rollout): claim FP8 weight updates only for FP8 weights
Sep 29, 2026
8f05bab
style(rollout): mark abort_request's broad except as deliberate
Sep 29, 2026
0ff2379
fix(rollout): release the barrier before giving up on any call
Sep 29, 2026
2f1e49a
fix(rollout): never route a collective_rpc reply on an unusable id
Sep 29, 2026
e534993
fix(rollout): reserve the names the engine's own protocol owns
Sep 29, 2026
04ea6b8
fix(rollout): serialize TP collective_rpc calls; advertise the RDMA l…
Sep 30, 2026
b754ade
test(rollout): pin that a queued reply outlives another rank's timeout
Oct 2, 2026
7c1f0dd
feat(rollout): receive trained weights over RCCL, transactionally
Sep 18, 2026
0e065f0
fix(rollout): verify reload coverage by the parameter actually written
Sep 28, 2026
55edb3b
fix(rollout): keep the RDMA stream in step at its edges; lift the fen…
Sep 30, 2026
0ce40f7
fix(rollout): fence a device fault found after commit; count only rea…
Sep 30, 2026
b591c68
fix(rollout): lift the fence only on a full reload; leave a refused s…
Oct 2, 2026
9c5728f
test(rollout): keep the IPC reload tests off torch's CUDA runtime
Oct 2, 2026
91c9d29
fix(rollout): free each bucket before the next; keep the full-load wa…
Oct 2, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
320 changes: 313 additions & 7 deletions atom/model_engine/async_proc.py

Large diffs are not rendered by default.

122 changes: 122 additions & 0 deletions atom/model_engine/capabilities.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,122 @@
# SPDX-License-Identifier: MIT
# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved.

"""What an engine and its workers support, so a caller can ask instead of guess.

LumenRL currently probes ATOM with ``hasattr(self.engine, ...)`` in four places
and otherwise hard-codes behaviour per backend. That works only while the two
repos move together; the moment they do not, a missing method is discovered as a
runtime failure mid-rollout rather than at negotiation time.

Two layers, answering different questions:

- :class:`WorkerCapabilities` -- what one rank can do. Collected from every rank
over ``collective_rpc``, so discovery uses the mechanism it describes.
- :class:`EngineCapabilities` -- the engine-wide answer: static topology plus the
**intersection** of the per-rank feature sets.

Intersection, not union, is the load-bearing choice. A feature present on some
ranks cannot be driven by a collective: the ranks that have it would block in
one collective while the rest went elsewhere, which deadlocks the group with no
error. Advertising it would be worse than not knowing.

Kept free of heavy imports for the same reason as ``collective_rpc``: the
dispatch and negotiation layers must import on a machine with no GPU build.
"""

from dataclasses import dataclass, field

# Bump when the collective-RPC wire contract changes incompatibly. A consumer
# that pins a version can refuse to negotiate rather than fail mid-collective.
COLLECTIVE_RPC_PROTOCOL_VERSION = 1


@dataclass(frozen=True)
class WorkerCapabilities:
"""One rank's view of itself."""

protocol_version: int
tp_rank: int
dp_rank_local: int
methods: frozenset[str] = frozenset()
features: frozenset[str] = frozenset()

@classmethod
def from_payload(cls, payload: object) -> "WorkerCapabilities":
"""Build from the dict a worker returned over the wire.

Tolerant of missing keys on purpose: an older worker that predates a
field should report less, not fail the whole negotiation.
"""
if not isinstance(payload, dict):
raise TypeError(
f"worker capabilities must be a dict, got {type(payload).__name__}"
)
return cls(
protocol_version=int(payload.get("protocol_version", 0)),
tp_rank=int(payload.get("tp_rank", -1)),
dp_rank_local=int(payload.get("dp_rank_local", 0)),
methods=frozenset(payload.get("methods", ())),
features=frozenset(payload.get("features", ())),
)


@dataclass(frozen=True)
class EngineCapabilities:
"""The engine-wide answer a consumer should negotiate against."""

protocol_version: int
tp_world_size: int
data_parallel_size: int
pipeline_parallel_size: int
kv_cache_dtype: str
# Only what every rank reports. See the module docstring.
methods: frozenset[str] = frozenset()
features: frozenset[str] = frozenset()
worker_count: int = 0
_mismatches: tuple[str, ...] = field(default=())

@classmethod
def from_workers(cls, *, config, workers) -> "EngineCapabilities":
if not workers:
raise ValueError("no worker capabilities to aggregate")

versions = {w.protocol_version for w in workers}
if len(versions) != 1:
raise RuntimeError(
f"workers disagree on the RPC protocol version: {sorted(versions)}"
)

methods = frozenset.intersection(*(w.methods for w in workers))
features = frozenset.intersection(*(w.features for w in workers))

# Record what was dropped, so a caller debugging a "missing" feature can
# see it was present-but-not-universal rather than absent everywhere.
union_methods = frozenset.union(*(w.methods for w in workers))
union_features = frozenset.union(*(w.features for w in workers))
mismatches = tuple(
sorted((union_methods - methods) | (union_features - features))
)

parallel = getattr(config, "parallel_config", None)
return cls(
protocol_version=versions.pop(),
tp_world_size=int(getattr(config, "tp_world_size", len(workers))),
data_parallel_size=int(getattr(parallel, "data_parallel_size", 1) or 1),
pipeline_parallel_size=int(
getattr(parallel, "pipeline_parallel_size", 1) or 1
),
kv_cache_dtype=str(getattr(config, "kv_cache_dtype", "auto")),
methods=methods,
features=features,
worker_count=len(workers),
_mismatches=mismatches,
)

def supports(self, name: str) -> bool:
"""Whether *name* is usable on every rank, as a method or a feature."""
return name in self.methods or name in self.features

def partial(self) -> tuple[str, ...]:
"""Names some ranks reported but not all, hence not advertised."""
return self._mismatches
110 changes: 110 additions & 0 deletions atom/model_engine/collective_rpc.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,110 @@
# SPDX-License-Identifier: MIT
# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved.

"""Wire types for the generic collective RPC.

Deliberately free of heavy imports. ``async_proc`` cannot be imported without a
real AITER build, and ``engine_utility`` needs these same types to build a
request, so keeping them here is what lets the dispatch layer stay importable on
a machine with no GPU -- which is also how the non-GPU test runner sees them.
"""

import queue
import threading
from collections.abc import Iterator
from contextlib import contextmanager
from dataclasses import dataclass

# The utility-command name the dispatch layer registers. Shared so the manager
# that sends it and the handler that receives it cannot drift apart.
COLLECTIVE_RPC_CMD = "collective_rpc"


@dataclass(frozen=True)
class RpcPayload:
"""One generic collective-RPC call: kwargs plus an explicit barrier request.

The plain worker wire format is ``(func_name, *args)``, which cannot carry
kwargs. Rather than change that format and every existing ``call_func``
caller, a generic call sends exactly one positional argument -- this -- and
the worker recognises it by type.

``request_id`` lets the manager match replies to the call that produced
them. The older KV channel matches by count instead, which is why it can
only ever have one aggregation outstanding.
"""

request_id: str
args: tuple = ()
kwargs: dict | None = None
barrier: bool = False

def call_kwargs(self) -> dict:
return self.kwargs or {}


@dataclass(frozen=True)
class RpcResult:
"""A generic-RPC reply, so a ``None`` return still reaches the caller.

``busy_loop`` only forwards non-``None`` worker returns, and
``call_func(wait_out=True)`` blocks on an untimed queue get, so a method
that legitimately returns ``None`` deadlocks the caller. Wrapping every
outcome -- including failures -- means the generic path always answers.
"""

request_id: str
tp_rank: int
value: object = None
error: str | None = None

@property
def ok(self) -> bool:
return self.error is None


class RpcResponseRouter:
"""Deliver DP-engine replies to whichever caller is waiting for them.

``broadcast_utility_command_sync`` reads a fixed number of replies off one
shared queue, so it matches by count: two overlapping callers take each
other's replies, and a late reply from an abandoned call is handed to the
next caller as its own. Upstream already recorded that biting the old
pull-based metrics.

Registering a request id gives that call its own queue, so concurrent calls
cannot collide and a reply nobody is waiting for is dropped rather than
mis-delivered.
"""

def __init__(self) -> None:
self._queues: dict[str, queue.Queue] = {}
self._lock = threading.Lock()

@contextmanager
def register(self, request_id: str) -> Iterator[queue.Queue]:
own: queue.Queue = queue.Queue()
with self._lock:
if request_id in self._queues:
raise RuntimeError(f"request id already in flight: {request_id}")
self._queues[request_id] = own
try:
yield own
finally:
# Unregister even on the timeout path, so a late reply is dropped by
# ``route`` rather than kept for a caller that has given up.
with self._lock:
self._queues.pop(request_id, None)

def route(self, request_id: str, item: object) -> bool:
"""Hand *item* to the waiter for *request_id*; False if nobody waits."""
with self._lock:
own = self._queues.get(request_id)
if own is None:
return False
own.put_nowait(item)
return True

def in_flight(self) -> int:
with self._lock:
return len(self._queues)
Loading
Loading