Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
7 changes: 5 additions & 2 deletions atom/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@
LayerQuantConfig,
get_quant_parser,
)
from atom.utils import envs, get_open_port
from atom.utils import envs
from atom.utils.distributed.utils import stateless_init_torch_distributed_process_group

if TYPE_CHECKING:
Expand Down Expand Up @@ -1018,7 +1018,10 @@ class ParallelConfig:
data_parallel_master_port: int = 29500
"""Port of the data parallel master."""

data_parallel_base_port: int = get_open_port()
data_parallel_base_port: int = 0
"""Model-runner rendezvous port. Zero requests an OS-assigned port locally."""
_managed_distributed_store: bool = field(default=False, init=False, repr=False)
"""CoreManager owns the store; every model runner connects as a client."""

data_parallel_master_ip: str = "127.0.0.1"

Expand Down
22 changes: 13 additions & 9 deletions atom/model_engine/block_table_codec.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
# SPDX-License-Identifier: MIT
# Copyright (C) 2024-2025, Advanced Micro Devices, Inc. All rights reserved.

"""Ship a forward's block tables to the TP workers as appends alone.
"""Ship a forward's block tables to the workers as appends alone.

Every forward RPC broadcasts one `ScheduledBatch` to every TP worker, and its
Every forward RPC broadcasts one `ScheduledBatch` to every worker, and its
`block_tables` are the bulk of it: one row per running request, the whole row
every step, growing one block per decode. At 50 seqs x 100k context that is
~313k ids -- 1.2 MiB pickled and unpickled per rank per step -- to say
Expand All @@ -25,7 +25,6 @@
as whole tables and resets both caches at once.
"""

import array
import copy
import logging
from dataclasses import dataclass
Expand Down Expand Up @@ -158,10 +157,10 @@ def _send_whole(self, batch, reason: str):


class BlockTableDeltaDecoder:
"""Worker side: rebuild `array("i")` rows from a `BlockTableDelta`."""
"""Worker side: rebuild versioned rows from a `BlockTableDelta`."""

def __init__(self):
self._rows: dict[int, array.array] = {}
self._rows: dict[int, BlockTable] = {}

def decode_rpc(self, func_name: str, args: list) -> list:
"""Decode `args[0]` in place if this is an encoded forward."""
Expand All @@ -184,8 +183,8 @@ def decode(self, batch):
f"{len(req_ids)} requests"
)

rows: list[array.array] = []
cached: dict[int, array.array] = {}
rows: list[BlockTable] = []
cached: dict[int, BlockTable] = {}
for i, req_id in enumerate(req_ids):
req_id = int(req_id)
base = int(delta.base_lengths[i])
Expand All @@ -203,9 +202,14 @@ def decode(self, batch):
# rewrite history the token processor may still be reading, so
# a row that grows is copied first. A row that did not grow is
# shared, which is the decode-step-with-no-new-block case.
row = previous if end == start else previous[:]
row = previous
if end != start:
row = BlockTable(previous)
# This is the next immutable snapshot in the same append
# lineage. Equal version + length still means equal ids.
row.version = previous.version
else:
row = array.array("i")
row = BlockTable()
if end > start:
row.frombytes(memoryview(delta.tail_values[start:end]).cast("B"))
rows.append(row)
Expand Down
43 changes: 37 additions & 6 deletions atom/model_engine/engine_core_mgr.py
Original file line number Diff line number Diff line change
Expand Up @@ -196,6 +196,7 @@ def _init_shared_state(
# is only safe to admit if it fits wherever the router sends it.
self.max_pool_tokens: int | None = None
self.engine_core_processes = []
self._distributed_stores = []
self.input_sockets = []
self.output_sockets = []
self.engine_core_identities = []
Expand Down Expand Up @@ -264,6 +265,35 @@ def _init_shared_state(
# scraping costs no round trip and cannot time out.
self.latest_metrics: dict[int, dict] = {}

def _start_distributed_store(
self, config: Config, *, multinode: bool = False
) -> None:
"""Bind the actual rendezvous server before publishing its port to workers."""
from torch.distributed import TCPStore

pc = config.parallel_config
if multinode and pc.data_parallel_base_port == 0:
raise ValueError(
"Multi-node DP requires the same nonzero --data-parallel-base-port "
"(or ATOM_DP_BASE_PORT) on every node."
)
if not multinode or pc.data_parallel_rank == 0:
store = TCPStore(
pc.data_parallel_master_ip,
pc.data_parallel_base_port,
is_master=True,
wait_for_workers=False,
)
self._distributed_stores.append(store)
pc.data_parallel_base_port = store.port
logger.info(
"%s: model-runner TCPStore listening on %s:%d",
self.label,
pc.data_parallel_master_ip,
store.port,
)
pc._managed_distributed_store = True

def __init__(self, config: Config):
pp_size = config.pipeline_parallel_size
self.pp_size = pp_size
Expand Down Expand Up @@ -344,6 +374,7 @@ def __init__(self, config: Config):
local_dp_ranks = []

try:
self._start_distributed_store(config, multinode=multinode)
for engine_index in range(self.local_engine_count):
assignment_index = engine_index // self.pp_size
dp_rank, local_dp_rank = rank_assignments[assignment_index]
Expand Down Expand Up @@ -803,6 +834,8 @@ def close(self):
except (ValueError, OSError):
pass

# Release the rendezvous server after stopping the local engine processes.
self._distributed_stores.clear()
logger.info(f"{self.label}: All EngineCores shut down")

def _send_request(self, dp_rank: int, payload: bytes) -> None:
Expand Down Expand Up @@ -1517,8 +1550,6 @@ def __init__(self, config: Config):
self._cu_shm = None

# Build per-process configs.
from atom.utils import get_open_port as _get_open_port

prefill_config = copy.deepcopy(config)
if config.disagg_prefill_max_num_seqs is not None:
prefill_config.max_num_seqs = config.disagg_prefill_max_num_seqs
Expand All @@ -1529,10 +1560,8 @@ def __init__(self, config: Config):
prefill_config.disagg_weight_ack_addr = weight_ack_addr
prefill_config.disagg_kvcache_ipc_addr = kvcache_ipc_addr
prefill_config.disagg_cu_shm_name = cu_shm_name
# Give prefill a distinct distributed rendezvous port so it doesn't
# collide with decode's data_parallel_base_port (both deep-copy the
# same port from config).
prefill_config.parallel_config.data_parallel_base_port = _get_open_port()
# Prefill gets its own server; decode honors the configured port.
prefill_config.parallel_config.data_parallel_base_port = 0

decode_config = copy.deepcopy(config)
decode_config.disagg_d2p_addr = d2p_addr
Expand Down Expand Up @@ -1613,6 +1642,8 @@ def _connect_proc(proc, in_addr, out_addr, ctrl_addr, name):
logger.info(f"{self.label}: {name} process started and connected")

try:
self._start_distributed_store(decode_config)
self._start_distributed_store(prefill_config)
# Start both processes simultaneously. Prefill binds the bootstrap
# PUSH socket and blocks on send() until decode connects and calls
# recv() — they rendezvous naturally without any sequential ordering.
Expand Down
Loading
Loading