Suppport TBO in ATOM - #515
Conversation
There was a problem hiding this comment.
Pull request overview
This PR adds TBO (Two-Batch Overlap) support to ATOM by introducing a micro-batch (ubatch) dual-thread execution path intended to overlap MoE communication with compute, and wires it into model execution and CUDAGraph capture.
Changes:
- Extend
ForwardContextto carry ubatch slicing info and add thread-local forward-context support for TBO worker threads. - Introduce a new
atom.utils.dbopackage implementing ubatch slicing, TBO thread/stream/event coordination, and aUBatchWrapperto run micro-batches in threads (including a TBO CUDAGraph capture path). - Integrate TBO into attention metadata builders, MORI prepare/finalize async flow, and model runner scheduling/DP synchronization + CLI/config flags.
Reviewed changes
Copilot reviewed 17 out of 17 changed files in this pull request and generated 1 comment.
Show a summary per file
| File | Description |
|---|---|
| atom/utils/forward_context.py | Add ubatch_slices to ForwardContext and thread-local context lookup for TBO threads. |
| atom/utils/dbo/ubatching.py | Implement TBOContext and module helpers for yield/stream switching and recv hooks. |
| atom/utils/dbo/ubatch_wrapper.py | Add UBatchWrapper to run ubatches in threads and optionally capture TBO execution in CUDAGraphs. |
| atom/utils/dbo/ubatch_splitting.py | Provide utilities to split batches into ubatches (decode and token-balanced prefill). |
| atom/utils/dbo/init.py | Export the new DBO/TBO public surface. |
| atom/models/deepseek_v2.py | Disable dual-stream MoE path when TBO is enabled. |
| atom/model_ops/moe.py | Enable MORI async/TBO wiring and instantiate per-ubatch MORI ops. |
| atom/model_ops/fused_moe/mori_prepare_finalize.py | Add async prepare/finalize paths (comm-stream + AsyncLL) and per-ubatch MORI op support. |
| atom/model_ops/fused_moe/modular_kernel.py | Route async prepare/finalize through TBO yield + hook mechanism; adjust finalize API usage. |
| atom/model_ops/attentions/backends.py | Add build_ubatch_prefill_metadata using split_attn_metadata. |
| atom/model_ops/attentions/aiter_mla.py | Allocate per-ubatch buffers and build per-ubatch attention metadata for TBO/CUDAGraph decode. |
| atom/model_ops/attentions/aiter_attention.py | Allocate per-ubatch buffers and build per-ubatch attention metadata for TBO/CUDAGraph decode. |
| atom/model_engine/scheduler.py | Extend next-batch info with request counts to support DP sync decisions. |
| atom/model_engine/model_runner.py | Wrap model with UBatchWrapper, create ubatch slices per batch, and add a TBO CUDAGraph capture path. |
| atom/model_engine/engine_core.py | Sync DP state including request counts; run dummy prefill with minimal reqs for TBO agreement. |
| atom/model_engine/arg_utils.py | Add CLI flags --enable-tbo and --low-latency. |
| atom/config.py | Add config toggles enable_tbo and enable_low_latency. |
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
0f6226d to
1072248
Compare
There was a problem hiding this comment.
Pull request overview
Copilot reviewed 17 out of 17 changed files in this pull request and generated 7 comments.
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
| config = get_current_atom_config() | ||
| if config is None: |
There was a problem hiding this comment.
tbo_enabled() calls get_current_atom_config(), which asserts if the config hasn’t been set; the subsequent if config is None: branch is therefore unreachable and may give a false sense that this helper is safe to call early. Consider either removing the None check, or switching to a non-asserting accessor (or try/except AssertionError) so tbo_enabled() reliably returns False when config isn’t initialized yet.
| config = get_current_atom_config() | |
| if config is None: | |
| try: | |
| config = get_current_atom_config() | |
| except AssertionError: |
| # All threads reach the barrier, then the main thread wakes thread 0 | ||
| self.ready_barrier.wait() | ||
|
|
||
| # Wait for our turn (thread 0 is woken by the main thread) | ||
| self.cpu_wait_event.wait() | ||
| self.cpu_wait_event.clear() | ||
|
|
There was a problem hiding this comment.
TBOContext.__enter__() blocks on ready_barrier.wait() and later on cpu_wait_event.wait() with no timeout or failure path. If any ubatch thread raises before reaching the barrier (e.g., device init / stream init failures) the main thread can block forever on its own ready_barrier.wait() in UBatchWrapper, resulting in a hard hang. Consider using a timeout + BrokenBarrierError handling, and/or ensuring threads always reach/abort the barrier in a finally block (e.g., ready_barrier.abort() plus waking any waiters) so errors propagate instead of deadlocking.
| try: | ||
| threads = [] | ||
| for i in range(N): | ||
| t = threading.Thread(target=_ubatch_thread, args=(i,)) | ||
| threads.append(t) | ||
| t.start() | ||
|
|
||
| self.ready_barrier.wait() | ||
| tbo_ctxs[0].cpu_wait_event.set() | ||
|
|
||
| for t in threads: | ||
| t.join() |
There was a problem hiding this comment.
_run_ubatches() waits on self.ready_barrier.wait() after starting worker threads. Because the barrier wait in TBOContext.__enter__() is unconditional and has no timeout, any exception in a worker before it reaches the barrier will cause the main thread to block forever here (and errors[idx] will never be checked). Consider adding timeouts + BrokenBarrierError handling, and aborting the barrier / setting the initial cpu_wait_event in the worker exception path so the system fails fast instead of deadlocking.
| if num_reqs >= self.max_num_seqs: | ||
| break | ||
| total_tokens += tokens | ||
| num_reqs += 1 |
There was a problem hiding this comment.
get_next_batch_info() can return (True, 0, 0) when self.waiting is non-empty but the first request would exceed max_num_batched_tokens (the loop breaks before adding any seq). This makes the DP sync state report “prefill with 0 tokens/reqs”, which can cause other ranks to compute global_max_tokens from decode ranks and run an oversized dummy prefill. Consider ensuring the function always reports at least one waiting request (e.g., include the first seq even if it exceeds the limit, or clamp total_tokens to the limit but keep num_reqs=1).
| num_reqs += 1 | |
| num_reqs += 1 | |
| if num_reqs == 0: | |
| first_seq = self.waiting[0] | |
| first_tokens = first_seq.num_tokens - first_seq.num_cached_tokens | |
| total_tokens = min(first_tokens, self.max_num_batched_tokens) | |
| num_reqs = 1 |
| def get_next_batch_info(self) -> tuple[bool, int, int]: | ||
| if self.waiting: | ||
| # new request is waiting, will do prefill | ||
| seq = self.waiting[0] | ||
| num_tokens = seq.num_tokens - seq.num_cached_tokens | ||
| return (True, num_tokens) | ||
| num_reqs = 0 |
There was a problem hiding this comment.
get_next_batch_info() now returns a 3-tuple, but the existing unit tests still expect the old (bool, int) shape (see tests/test_scheduler.py:323-336). Update/add tests to validate the new (is_prefill, total_tokens, num_reqs) contract for empty/waiting/running cases (and any edge cases introduced by the new batching loop).
1072248 to
21b933a
Compare
There was a problem hiding this comment.
Pull request overview
Copilot reviewed 17 out of 17 changed files in this pull request and generated 5 comments.
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
| padding = torch.empty(padding_shape, dtype=x.dtype, device=x.device) | ||
| padding.zero_() | ||
| # padding.zero_() | ||
| padded_x = torch.cat([x, padding], dim=dim) |
There was a problem hiding this comment.
Padding is allocated with torch.empty and left uninitialized before being all-gathered across DP ranks. This can leak uninitialized GPU memory contents across ranks and introduce non-determinism (and potentially NaNs) into the MoE routing/compute for the padded tokens. Prefer initializing padding (e.g., zeros) before concatenation.
| def get_next_batch_info(self) -> tuple[bool, int, int]: | ||
| if self.waiting: | ||
| # new request is waiting, will do prefill | ||
| seq = self.waiting[0] | ||
| num_tokens = seq.num_tokens - seq.num_cached_tokens | ||
| return (True, num_tokens) | ||
| num_reqs = 0 | ||
| total_tokens = 0 | ||
| for seq in self.waiting: | ||
| tokens = seq.num_tokens - seq.num_cached_tokens | ||
| if total_tokens + tokens > self.max_num_batched_tokens: | ||
| break | ||
| if num_reqs >= self.max_num_seqs: | ||
| break | ||
| total_tokens += tokens | ||
| num_reqs += 1 | ||
| return (True, total_tokens, num_reqs) | ||
| elif self.running: | ||
| # decode | ||
| num_tokens = len(self.running) | ||
| return (False, num_tokens) | ||
| return (False, num_tokens, num_tokens) | ||
| else: | ||
| # No requests | ||
| return (False, 0) | ||
| return (False, 0, 0) |
There was a problem hiding this comment.
get_next_batch_info now returns a 3-tuple (is_prefill, num_tokens, num_reqs), which breaks existing unit tests that assert a 2-tuple return (e.g. tests/test_scheduler.py around the TestGetNextBatchInfo section). Please update/add tests to cover the new return shape and the new prefill counting logic.
| # Build per-ubatch ForwardContexts from pre-allocated forward_vars. | ||
| full_graph_bs = ctx.context.graph_bs | ||
| # only padding for all_gather/reduce_scatter pass | ||
| all_gahter_dp_size = self._get_dp_size() if self.dp_gather_scatter else 1 |
There was a problem hiding this comment.
Variable name all_gahter_dp_size appears to be misspelled, which makes the code harder to scan/search (especially near other all-gather logic). Rename it to all_gather_dp_size (or similar) for clarity.
| gpu_comm_done_events = [torch.Event() for _ in range(num_micro_batches)] | ||
| gpu_compute_done_events = [torch.Event() for _ in range(num_micro_batches)] |
There was a problem hiding this comment.
GPU events are created with torch.Event(), but this code later uses CUDA stream wait_event/record semantics and the rest of the codebase uses torch.cuda.Event(). torch.Event() is not a CUDA event and may not exist / may not work here; use torch.cuda.Event() (optionally with blocking=False / enable_timing=False as needed).
| gpu_comm_done_events = [torch.Event() for _ in range(num_micro_batches)] | |
| gpu_compute_done_events = [torch.Event() for _ in range(num_micro_batches)] | |
| gpu_comm_done_events = [torch.cuda.Event() for _ in range(num_micro_batches)] | |
| gpu_compute_done_events = [torch.cuda.Event() for _ in range(num_micro_batches)] |
There was a problem hiding this comment.
Pull request overview
Copilot reviewed 21 out of 21 changed files in this pull request and generated 4 comments.
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
| ```bash | ||
| --enable-tbo # high-throughput mode (default) | ||
| --enable-tbo high-throughput # explicit high-throughput mode | ||
| --enable-tbo low-latency # AsyncLL MORI kernel for low-latency overlap | ||
| ``` |
There was a problem hiding this comment.
The documented CLI suggests --enable-tbo takes a mode argument (high-throughput/low-latency), but the actual CLI parsing (see atom/model_engine/arg_utils.py) defines --enable-tbo as a boolean flag and uses a separate --low-latency flag. Update this guide to match the implemented flags, or update the CLI to accept the documented values.
| # only padding for all_gather/reduce_scatter pass | ||
| all_gahter_dp_size = self._get_dp_size() if self.dp_gather_scatter else 1 | ||
| forward_contexts = [] | ||
| ub_inputs = [] | ||
| for i, ub_slice in enumerate(ctx.ubatch_slices): | ||
| if i < N - 1: | ||
| padded_bs = full_graph_bs // N | ||
| else: | ||
| padded_bs = full_graph_bs - (full_graph_bs // N) * (N - 1) | ||
| ub_ctx = self._make_ubatch_context( | ||
| ctx, | ||
| ub_slice, | ||
| padded_bs, | ||
| i, | ||
| ub_graph_bs=padded_bs * all_gahter_dp_size, |
There was a problem hiding this comment.
Typo in variable name: all_gahter_dp_size should be all_gather_dp_size (or similar). Fixing this improves readability and avoids propagating the misspelling.
f445523 to
99418eb
Compare
There was a problem hiding this comment.
Pull request overview
Copilot reviewed 21 out of 21 changed files in this pull request and generated 3 comments.
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
| # only padding for all_gather/reduce_scatter pass | ||
| all_gahter_dp_size = self._get_dp_size() if self.dp_gather_scatter else 1 | ||
| forward_contexts = [] | ||
| ub_inputs = [] | ||
| for i, ub_slice in enumerate(ctx.ubatch_slices): | ||
| if i < N - 1: | ||
| padded_bs = full_graph_bs // N | ||
| else: | ||
| padded_bs = full_graph_bs - (full_graph_bs // N) * (N - 1) | ||
| ub_ctx = self._make_ubatch_context( | ||
| ctx, | ||
| ub_slice, | ||
| padded_bs, | ||
| i, | ||
| ub_graph_bs=padded_bs * all_gahter_dp_size, | ||
| ) |
There was a problem hiding this comment.
Variable name all_gahter_dp_size is misspelled (should be all_gather_dp_size). This typo makes the intent harder to follow and is easy to propagate to other code paths.
| def maybe_create_ubatch_slices( | ||
| num_reqs: int, | ||
| num_tokens: int, | ||
| num_ubatches: int = 2, | ||
| is_prefill: bool = False, | ||
| num_scheduled_tokens: Optional[list[int]] = None, | ||
| max_tokens_per_ubatch: Optional[int] = None, | ||
| ) -> Optional[list[UBatchSlice]]: | ||
| """Split a batch into N micro-batch slices. | ||
|
|
||
| For decode: split by request count (uniform tokens per request). | ||
| For prefill: token-balanced split so each ubatch has roughly equal | ||
| token count, respecting request boundaries. | ||
|
|
||
| Returns None if the batch is too small to split or if the split | ||
| would produce a ubatch exceeding max_tokens_per_ubatch. | ||
| """ | ||
| if num_ubatches <= 1: | ||
| return None | ||
|
|
||
| if num_reqs < num_ubatches: | ||
| return None | ||
|
|
||
| if num_scheduled_tokens is not None: | ||
| # Prefill: token-balanced split | ||
| return _split_prefill_balanced( | ||
| num_reqs, | ||
| num_scheduled_tokens, | ||
| num_ubatches, | ||
| max_tokens_per_ubatch, | ||
| ) |
There was a problem hiding this comment.
maybe_create_ubatch_slices advertises splitting into num_ubatches micro-batches, but the prefill path (num_scheduled_tokens is not None) always returns exactly 2 slices via _split_prefill_balanced, ignoring num_ubatches. Either enforce num_ubatches == 2 for prefill (validate and return None/raise) or implement true N-way token-balanced splitting to match the API contract/docstring.
| def maybe_create_ubatch_slices( | ||
| num_reqs: int, | ||
| num_tokens: int, | ||
| num_ubatches: int = 2, | ||
| is_prefill: bool = False, | ||
| num_scheduled_tokens: Optional[list[int]] = None, | ||
| max_tokens_per_ubatch: Optional[int] = None, | ||
| ) -> Optional[list[UBatchSlice]]: | ||
| """Split a batch into N micro-batch slices. | ||
|
|
||
| For decode: split by request count (uniform tokens per request). | ||
| For prefill: token-balanced split so each ubatch has roughly equal | ||
| token count, respecting request boundaries. | ||
|
|
||
| Returns None if the batch is too small to split or if the split | ||
| would produce a ubatch exceeding max_tokens_per_ubatch. | ||
| """ | ||
| if num_ubatches <= 1: | ||
| return None | ||
|
|
||
| if num_reqs < num_ubatches: | ||
| return None | ||
|
|
||
| if num_scheduled_tokens is not None: | ||
| # Prefill: token-balanced split | ||
| return _split_prefill_balanced( | ||
| num_reqs, | ||
| num_scheduled_tokens, | ||
| num_ubatches, | ||
| max_tokens_per_ubatch, | ||
| ) | ||
|
|
||
| # Decode: uniform split by request count | ||
| slices = [] | ||
| reqs_per_ub = num_reqs // num_ubatches | ||
| tokens_per_req = num_tokens // num_reqs if num_reqs > 0 else 1 | ||
|
|
||
| for i in range(num_ubatches): | ||
| req_start = i * reqs_per_ub | ||
| req_end = (i + 1) * reqs_per_ub if i < num_ubatches - 1 else num_reqs | ||
| tok_start = req_start * tokens_per_req | ||
| tok_end = req_end * tokens_per_req | ||
| slices.append( | ||
| UBatchSlice( | ||
| request_slice=slice(req_start, req_end), | ||
| token_slice=slice(tok_start, tok_end), | ||
| ) | ||
| ) | ||
| return slices |
There was a problem hiding this comment.
No unit tests were added for the new ubatch-splitting logic (decode uniform split, prefill token-balanced split, and max_tokens_per_ubatch rejection). Given the edge cases (non-uniform token distributions, small batches, boundary conditions), adding targeted tests would help prevent regressions.
54e7d15 to
6c2e807
Compare
There was a problem hiding this comment.
Pull request overview
Copilot reviewed 20 out of 20 changed files in this pull request and generated 6 comments.
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
| # Barrier: ubatch threads + main thread | ||
| self.ready_barrier = threading.Barrier(3) # 2 ubatch threads + 1 main | ||
| # TBO CUDAGraph storage: keyed by (graph_bs, max_q_len) |
There was a problem hiding this comment.
UBatchWrapper hard-codes ready_barrier=threading.Barrier(3) (2 ubatch threads + main). But _run_ubatches() uses N = len(ctx.ubatch_slices) and starts N threads. If ubatch_slices is ever created with N != 2, this will deadlock at the barrier. Either assert N == 2 (since TBO is specifically two-batch overlap) or construct/reset the barrier dynamically as threading.Barrier(N + 1).
| # only padding for all_gather/reduce_scatter pass | ||
| all_gahter_dp_size = self._get_dp_size() if self.dp_gather_scatter else 1 | ||
| forward_contexts = [] |
There was a problem hiding this comment.
Typo in variable name: all_gahter_dp_size. This makes the code harder to read/grep and is easy to propagate. Rename to all_gather_dp_size.
| def get_next_batch_info(self) -> tuple[bool, int, int]: | ||
| if self.waiting: | ||
| # new request is waiting, will do prefill | ||
| seq = self.waiting[0] | ||
| num_tokens = seq.num_tokens - seq.num_cached_tokens | ||
| return (True, num_tokens) | ||
| num_reqs = 0 | ||
| total_tokens = 0 | ||
| for seq in self.waiting: | ||
| tokens = seq.num_tokens - seq.num_cached_tokens | ||
| if total_tokens + tokens > self.max_num_batched_tokens: | ||
| break | ||
| if num_reqs >= self.max_num_seqs: | ||
| break | ||
| total_tokens += tokens | ||
| num_reqs += 1 | ||
| return (True, total_tokens, num_reqs) |
There was a problem hiding this comment.
get_next_batch_info() now returns a 3-tuple (is_prefill, total_tokens, num_reqs). This is a breaking change for callers/tests that still expect the previous 2-tuple. Please update any remaining call sites (e.g., tests/test_scheduler.py) and consider adding a brief docstring comment about the new third value to avoid silent misuse.
| dummy_reqs = min( | ||
| global_max_reqs, 2 | ||
| ) # dummy reqs at 2: just enough for TBO agreement, avoid wasting compute. | ||
| logger.info( | ||
| f"{self.label}: Running dummy prefill ({global_max_tokens} tokens) " | ||
| f"{self.label}: Running dummy prefill ({global_max_tokens} tokens, {dummy_reqs} reqs) " | ||
| f"to sync with other DP ranks doing prefill" | ||
| ) | ||
| self._execute_dummy_prefill(global_max_tokens) | ||
| self._execute_dummy_prefill(global_max_tokens, dummy_reqs) |
There was a problem hiding this comment.
dummy_reqs can become 0 when global_max_reqs is 0, which makes the log message inaccurate and passes 0 into dummy_prefill_execution (even though the callee clamps it back to >= 1). Clamp dummy_reqs to at least 1 here (before logging/calling) so behavior and logs are consistent.
|
|
||
|
|
||
| def tbo_overlap_enabled() -> bool: | ||
| return False |
There was a problem hiding this comment.
tbo_overlap_enabled() currently always returns False, but modular_kernel uses it in assertions to guard the non-async path. With the current implementation, those assertions can never trigger, and the function name becomes misleading. Consider implementing this as return tbo_active() (or removing it entirely and calling tbo_active/tbo_enabled directly) so the guard actually reflects whether TBO overlap is active in the current thread.
| return False | |
| return tbo_active() |
| padding = torch.empty(padding_shape, dtype=x.dtype, device=x.device) | ||
| padding.zero_() | ||
| # padding.zero_() | ||
| padded_x = torch.cat([x, padding], dim=dim) |
There was a problem hiding this comment.
pad_for_all_gather() builds padding via torch.empty() and no longer zero-fills it (padding.zero_() is commented out). This can introduce nondeterministic / garbage values into the gathered tensor, which can corrupt router logits / MoE routing (and generally produce incorrect outputs) even if the extra tokens are later unpadded. Prefer zero-initializing the padding (or using torch.zeros for the padding tensor) so padded tokens are neutral.
Motivation
We enable TBO with dp attn + mori
--enable-dp-attention --enable-expert-parallel --enable-tbo
GPT-OSS:
1. DP + EP mori + TBO(prefill only) vs DP + EP mori
Config:
-tp 2 --enable-dp-attention --enable-expert-parallel --enable-tbo256 BS
512 BS
2. DP + all_gather/reduce_scatter + TBO (prefill only) vs DP + all_gather/reduce_scatter
Config:
-tp 2 --enable-dp-attention --enable-tbo256 BS
512 BS
Deepseek:
DP + EP MORI + TBO + MTP 3 (Speculative Decoding)
Overlap:
perfill:

decode:

Technical Details
Test Plan
Test Result
Submission Checklist