Skip to content

Suppport TBO in ATOM - #515

Merged
valarLip merged 15 commits into
mainfrom
zlr/tbo_dev
Apr 16, 2026
Merged

valarLip merged 15 commits into
mainfrom
zlr/tbo_dev

Conversation

@ZhangLirong-amd

@ZhangLirong-amd ZhangLirong-amd commented Apr 8, 2026 •

Copy link
Copy Markdown
Collaborator

Motivation

We enable TBO with dp attn + mori
--enable-dp-attention --enable-expert-parallel --enable-tbo

  • --enable-tbo → enable_tbo prefill only
  • --enable-tbo all → enable prefill + decode
  • --all2all-backend low-latency → enable low_latency
  • --all2all-backend or --all2all-backend high-throughput → enable high-throughput

GPT-OSS:

MORI_SHMEM_MODE=ISOLATION python3 -m atom.entrypoints.openai_server --model /data/models/openai/gpt-oss-120b/ -tp 2 --port 5678  --gpu-memory-utilization 0.4  --enable-dp-attention   --server-port 7777 --torch-profiler-dir ./log --enable-expert-parallel --enable-tbo

1. DP + EP mori + TBO(prefill only) vs DP + EP mori

Config: -tp 2 --enable-dp-attention --enable-expert-parallel --enable-tbo

256 BS

Configuration Duration (s) Req Throughput (req/s) Output Throughput (tok/s) Total Throughput (tok/s)
dp + ep mori 19.19 13.34 13,660 27,320
dp + ep mori + tbo (all) 24.53 10.44 10,688 21,376
dp + ep mori + tbo (prefill only) 18.45 13.87 14,206 28,411
TBO prefill only Impact -3.85% +3.97% +4.00% +3.99%

512 BS

Configuration Duration (s) Req Throughput (req/s) Output Throughput (tok/s) Total Throughput (tok/s)
dp + ep mori 33.94 15.09 15,449 30,898
dp + ep mori + tbo (all) 37.16 13.78 14,109 28,219
dp + ep mori + tbo (prefill only) 29.02 17.64 18,067 36,134
TBO prefill only Impact -14.50% +16.90% +16.94% +16.94%

2. DP + all_gather/reduce_scatter + TBO (prefill only) vs DP + all_gather/reduce_scatter

TORCH_NCCL_BLOCKING_WAIT=1 python3 -m atom.entrypoints.openai_server --model /data/models/openai/gpt-oss-120b/ -tp 2 --port 5678  --gpu-memory-utilization 0.4  --enable-dp-attention   --server-port 7777 --torch-profiler-dir ./log --enable-tbo

Config: -tp 2 --enable-dp-attention --enable-tbo

256 BS

Configuration Duration (s) Req Throughput (req/s) Output Throughput (tok/s) Total Throughput (tok/s)
dp 17.40 14.71 15,066 30,132
dp + tbo (all) 22.38 11.44 11,714 23,427
dp + tbo (prefill only) 17.04 15.02 15,380 30,759
TBO prefill only Impact -2.07% +2.11% +2.08% +2.08%

512 BS

Configuration Duration (s) Req Throughput (req/s) Output Throughput (tok/s) Total Throughput (tok/s)
dp 29.86 17.15 17,560 35,120
dp + tbo (all) 35.19 14.55 14,898 29,796
dp + tbo (prefill only) 26.29 19.48 19,944 39,889
TBO prefill only Impact -11.96% +13.59% +13.57% +13.58%

Deepseek:

MORI_SHMEM_MODE=ISOLATION python3 -m atom.entrypoints.openai_server --model /data/deepseek-ai/DeepSeek-R1-0528/ -tp 8 --port 5678  --gpu-memory-utilization 0.4  --enable-dp-attention --enable-expert-parallel --enable-tbo --server-port 7777

DP + EP MORI + TBO + MTP 3 (Speculative Decoding)

Concurrency Configuration Output Throughput (tok/s) Total Throughput (tok/s) TBO Impact
256 dp + ep mori + mtp3 6,375 12,764 —
256 dp + ep mori + tbo + mtp3 6,590 13,196 +3.4%

Overlap:

perfill:
image

decode:
image

Technical Details

Test Plan

Test Result

Submission Checklist

Copilot AI review requested due to automatic review settings April 8, 2026 05:38

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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 ForwardContext to carry ubatch slicing info and add thread-local forward-context support for TBO worker threads.
  • Introduce a new atom.utils.dbo package implementing ubatch slicing, TBO thread/stream/event coordination, and a UBatchWrapper to 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.

Comment thread atom/model_engine/model_runner.py Outdated
Copilot AI review requested due to automatic review settings April 8, 2026 09:26

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment on lines +17 to +18
config = get_current_atom_config()
if config is None:

Copilot AI Apr 8, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Suggested change
config = get_current_atom_config()
if config is None:
try:
config = get_current_atom_config()
except AssertionError:

Copilot uses AI. Check for mistakes.
Comment thread atom/utils/tbo/ubatching.py
Comment on lines +84 to +90
# 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()

Copilot AI Apr 8, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copilot uses AI. Check for mistakes.
Comment on lines +145 to +156
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()

Copilot AI Apr 8, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

_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.

Copilot uses AI. Check for mistakes.
if num_reqs >= self.max_num_seqs:
break
total_tokens += tokens
num_reqs += 1

Copilot AI Apr 8, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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).

Suggested change
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

Copilot uses AI. Check for mistakes.
Comment thread atom/model_ops/fused_moe/modular_kernel.py
Comment on lines +587 to +590
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

Copilot AI Apr 8, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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).

Copilot uses AI. Check for mistakes.
Copilot AI review requested due to automatic review settings April 15, 2026 03:37

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment thread atom/model_ops/moe.py
Comment on lines 191 to 193
padding = torch.empty(padding_shape, dtype=x.dtype, device=x.device)
padding.zero_()
# padding.zero_()
padded_x = torch.cat([x, padding], dim=dim)

Copilot AI Apr 15, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copilot uses AI. Check for mistakes.
Comment on lines +587 to +607
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)

Copilot AI Apr 15, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copilot uses AI. Check for mistakes.
# 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

Copilot AI Apr 15, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copilot uses AI. Check for mistakes.
Comment thread atom/models/deepseek_v2.py
Comment on lines +303 to +304
gpu_comm_done_events = [torch.Event() for _ in range(num_micro_batches)]
gpu_compute_done_events = [torch.Event() for _ in range(num_micro_batches)]

Copilot AI Apr 15, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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).

Suggested change
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)]

Copilot uses AI. Check for mistakes.
Copilot AI review requested due to automatic review settings April 15, 2026 06:16

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment thread recipes/TBO.md
Comment on lines +28 to +32
```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
```

Copilot AI Apr 15, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copilot uses AI. Check for mistakes.
Comment thread atom/model_engine/arg_utils.py
Comment on lines +210 to +224
# 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,

Copilot AI Apr 15, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copilot uses AI. Check for mistakes.
Comment thread atom/models/deepseek_v2.py
Copilot AI review requested due to automatic review settings April 15, 2026 12:25

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment on lines +210 to +225
# 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,
)

Copilot AI Apr 15, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copilot uses AI. Check for mistakes.
Comment on lines +23 to +53
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,
)

Copilot AI Apr 15, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copilot uses AI. Check for mistakes.
Comment on lines +23 to +71
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

Copilot AI Apr 15, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copilot uses AI. Check for mistakes.
Copilot AI review requested due to automatic review settings April 16, 2026 02:06

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment on lines +56 to +58
# 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)

Copilot AI Apr 16, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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).

Copilot uses AI. Check for mistakes.
Comment on lines +210 to +212
# 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 = []

Copilot AI Apr 16, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copilot uses AI. Check for mistakes.
Comment on lines +587 to +600
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)

Copilot AI Apr 16, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copilot uses AI. Check for mistakes.
Comment on lines +389 to +396
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)

Copilot AI Apr 16, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copilot uses AI. Check for mistakes.


def tbo_overlap_enabled() -> bool:
return False

Copilot AI Apr 16, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Suggested change
return False
return tbo_active()

Copilot uses AI. Check for mistakes.
Comment thread atom/model_ops/moe.py
Comment on lines 196 to 198
padding = torch.empty(padding_shape, dtype=x.dtype, device=x.device)
padding.zero_()
# padding.zero_()
padded_x = torch.cat([x, padding], dim=dim)

Copilot AI Apr 16, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copilot uses AI. Check for mistakes.
@valarLip
valarLip merged commit a022e6c into main Apr 16, 2026
33 of 41 checks passed
@valarLip
valarLip deleted the zlr/tbo_dev branch April 16, 2026 13:40
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants