diff --git a/python/sglang/srt/batch_overlap/operations_strategy.py b/python/sglang/srt/batch_overlap/operations_strategy.py index d39ad8385..af2171e3d 100644 --- a/python/sglang/srt/batch_overlap/operations_strategy.py +++ b/python/sglang/srt/batch_overlap/operations_strategy.py @@ -63,6 +63,15 @@ class OperationsStrategy: for layer in layers ] ) + elif layer_name == "DeepseekV4DecoderLayer": + return OperationsStrategy.concat( + [ + _compute_moe_deepseek_v4_layer_operations_strategy_tbo( + layer, forward_mode + ) + for layer in layers + ] + ) else: raise NotImplementedError @@ -150,6 +159,73 @@ def _compute_moe_deepseek_blog_decode(layer): ) +# -------------------------------- Strategy for DeepSeek V4 --------------------------------------- + + +# DSV4 prefill TBO (EP / mori path). Cross-layer mHC fusion is disabled under +# TBO, so each layer is self-contained: attn-side mHC pre+norm -> attn -> +# ffn-side mHC pre+norm -> MoE (a2a dispatch/combine overlapped) -> mHC post. +# The MoE ops are reused from self.mlp (DeepseekV2MoE) and decompose +# forward_deepep; the layer-level op_mhc_* wrap DSV4's hc_pre / hc_post. +def _compute_moe_deepseek_v4_layer_operations_strategy_tbo( + layer: torch.nn.Module, + forward_mode: ForwardMode, +) -> OperationsStrategy: + if forward_mode == ForwardMode.EXTEND: + return _compute_moe_deepseek_v4_prefill(layer) + else: + # Decode TBO for DSV4 is not implemented yet (ATOM data: decode TBO + # regresses; needs cuda-graph capture work). Prefill-only for now. + raise NotImplementedError( + f"DeepseekV4 TBO only supports prefill (EXTEND), got {forward_mode=}" + ) + + +def _compute_moe_deepseek_v4_prefill(layer): + from sglang.srt.layers.moe import get_moe_a2a_backend + + if get_moe_a2a_backend().is_none(): + # Non-EP DP TP-MoE: overlap the DP all_gatherv (gather) + reduce_scatterv + # (combine) with the other ubatch's attn+MoE compute (ATOM's DSV4 path). + ops = [ + layer.op_mhc_prepare_attn, + layer.self_attn.op_attn, + layer.op_mhc_post_attn_pre_mlp, + layer.op_gather_a, + operations.YieldOperation(), + layer.op_gather_b, + layer.op_moe, + layer.op_combine_a, + operations.YieldOperation(), + layer.op_combine_b, + layer.op_mhc_postprocess, + ] + else: + # EP / mori a2a: reuse DeepseekV2MoE's deepep dispatch/combine ops. + ops = [ + layer.op_mhc_prepare_attn, + layer.self_attn.op_attn, + layer.op_mhc_post_attn_pre_mlp, + layer.mlp.op_gate, + layer.mlp.op_select_experts, + layer.mlp.op_dispatch_a, + operations.YieldOperation(), + layer.mlp.op_dispatch_b, + layer.mlp.op_experts, + layer.mlp.op_combine_a, + operations.YieldOperation(), + layer.mlp.op_shared_experts, + layer.mlp.op_combine_b, + layer.mlp.op_output, + layer.op_mhc_postprocess, + ] + return OperationsStrategy( + deep_gemm_num_sms=None, + tbo_delta_stages=0, + operations=ops, + ) + + # -------------------------------- Strategy for Qwen3 --------------------------------------- diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index 7d67a71ec..c4a33f091 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -103,9 +103,28 @@ def _is_two_chunk_split_enabled(extend_lens: Sequence[int]) -> bool: overall_sum = sum(extend_lens) threshold = get_tbo_token_distribution_threshold() assert threshold <= 0.5, f"{threshold=}" - return left_sum < overall_sum * threshold or left_sum > overall_sum * ( + want_two_chunk = left_sum < overall_sum * threshold or left_sum > overall_sum * ( 1 - threshold ) + if not want_two_chunk: + return False + + # Two-chunk splits a single seq across both micro-batches by cutting at + # overall_sum // 2. child_a then spans seqs [0 : split_seq_index + 1] + # (batch_size = split_seq_index + 1) but only receives overall_sum // 2 + # query tokens. For a degenerate batch (a single seq, or a near-empty + # DP-sync batch) this cut is 0 or tiny, leaving child_a with more seqs + # than query tokens (e.g. (bs=1, tok=0)). That violates the DSV4 compress + # planner invariant `batch_size <= num_q_tokens` and crashes the kernel. + # Fall back to a seq-boundary split, whose child_a is seq-aligned (each + # seq contributes >= 1 token) and cannot become empty-with-count. + split_seq_index = _split_array_by_cum_less_than_half(extend_lens) + child_a_batch_size = split_seq_index + 1 + child_a_num_q_tokens = overall_sum // 2 + if child_a_batch_size > child_a_num_q_tokens: + return False + + return True def _split_extend_seqs(arr: Sequence[int]) -> int: @@ -652,9 +671,10 @@ class TboForwardBatchPreparer: output_dict[key] = old_value[start_token_index:end_token_index] attention_tp_size = get_parallel().attn_tp_size - output_dict["tbo_padded_len"] = ( + _tbo_padded_len = ( (end_token_index - start_token_index - 1) // attention_tp_size + 1 ) * attention_tp_size + output_dict["tbo_padded_len"] = _tbo_padded_len for key in [ "req_pool_indices", diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py index 152426276..ec225dd17 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py @@ -406,6 +406,14 @@ class _GraphBucket(enum.Enum): class DeepseekV4HipRadixBackend( AttentionBackend, C4IndexerBackendMixin, CompressorBackendMixin ): + # DSV4 TBO runs ONLY in eager prefill (prefill cuda-graph is disabled); + # decode/target-verify graphs are non-TBO (primary backend only). So the TBO + # child backends must not be driven through cuda-graph capture/replay — doing + # so rebuilds this backend's compressor/indexer metadata per replay step on + # both children and leaks ROCm HSA resources (HSA_STATUS_ERROR_OUT_OF_RESOURCES). + # TboAttnBackend reads this to skip children in the *_graph paths only. + tbo_supports_cuda_graph = False + def __init__( self, model_runner: ModelRunner, diff --git a/python/sglang/srt/layers/attention/tbo_backend.py b/python/sglang/srt/layers/attention/tbo_backend.py index 335d8cc4e..d595f3520 100644 --- a/python/sglang/srt/layers/attention/tbo_backend.py +++ b/python/sglang/srt/layers/attention/tbo_backend.py @@ -25,6 +25,20 @@ class TboAttnBackend(AttentionBackend): children=[creator() for _ in range(2)], ) + def _children_use_cuda_graph(self) -> bool: + """Whether the TBO child backends participate in CUDA-graph capture/replay. + + Some models only run TBO in eager prefill and keep their graph-captured + modes (decode / target-verify) NON-TBO on the primary backend. For those, + the children must NOT be driven through the cuda-graph paths: doing so + rebuilds their per-step metadata on every replay even though the captured + graph never uses them. For DeepSeek-V4 that metadata build (compressor / + indexer) leaks ROCm HSA resources across the 2 children -> eventual + HSA_STATUS_ERROR_OUT_OF_RESOURCES. Eager prefill TBO (init_forward_metadata) + is unaffected; only the *_graph paths are gated. + """ + return getattr(self.primary, "tbo_supports_cuda_graph", True) + def init_forward_metadata_out_graph( self, forward_batch: "ForwardBatch", @@ -33,6 +47,8 @@ class TboAttnBackend(AttentionBackend): self.primary.init_forward_metadata_out_graph( forward_batch=forward_batch, in_capture=in_capture ) + if not self._children_use_cuda_graph(): + return tbo_children = getattr(forward_batch, "tbo_children", None) if tbo_children is not None: for child, forward_batch_child in zip( @@ -97,6 +113,8 @@ class TboAttnBackend(AttentionBackend): def init_forward_metadata_in_graph(self, forward_batch: "ForwardBatch"): self.primary.init_forward_metadata_in_graph(forward_batch=forward_batch) + if not self._children_use_cuda_graph(): + return tbo_children = getattr(forward_batch, "tbo_children", None) if tbo_children is not None: for child, forward_batch_child in zip( @@ -118,17 +136,23 @@ class TboAttnBackend(AttentionBackend): def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): self.primary.init_cuda_graph_state(max_bs=max_bs, max_num_tokens=max_num_tokens) + if not self._children_use_cuda_graph(): + return for item in self.children: # TODO for children, maybe can provide *smaller* max_bs to optimize item.init_cuda_graph_state(max_bs=max_bs, max_num_tokens=max_num_tokens) def on_after_cuda_graph_warmup(self): self.primary.on_after_cuda_graph_warmup() + if not self._children_use_cuda_graph(): + return for child in self.children: child.on_after_cuda_graph_warmup() def get_cuda_graph_seq_len_fill_value(self): ans = self.primary.get_cuda_graph_seq_len_fill_value() + if not self._children_use_cuda_graph(): + return ans for child in self.children: assert ans == child.get_cuda_graph_seq_len_fill_value() return ans @@ -145,6 +169,19 @@ class TboAttnBackend(AttentionBackend): def get_indexer_metadata(self, layer_id: int, forward_batch: "ForwardBatch"): return self.primary.get_indexer_metadata(layer_id, forward_batch) + def __getattr__(self, name): + # Delegate backend-specific attributes/methods not explicitly wrapped + # above (e.g. DSV4's get_unified_swa_loc / get_swa_out_cache_loc, which + # the model calls directly via get_attn_backend()) to the primary + # full-batch backend. Inside TBO the per-child backend is resolved + # directly from the forward context, so this path only serves the + # non-overlapped forward (warmup / decode / TBO-ineligible batches). + # NOTE: __getattr__ runs only when normal lookup fails; guard `primary` + # to avoid infinite recursion before __init__ sets it. + if name == "primary": + raise AttributeError(name) + return getattr(self.primary, name) + def _build_tbo_child_replay_fb_view( fb_view, diff --git a/python/sglang/srt/layers/dp_attention.py b/python/sglang/srt/layers/dp_attention.py index 00495d84a..ce535fa7a 100644 --- a/python/sglang/srt/layers/dp_attention.py +++ b/python/sglang/srt/layers/dp_attention.py @@ -691,6 +691,112 @@ def dp_reduce_scatter_tensor(output: torch.Tensor, input: torch.Tensor): get_attention_tp_group().all_gather_into_tensor(output, scattered_local_tokens) +# --------------------------------------------------------------------------- +# Two-batch-overlap (non-EP / DP TP-MoE) async gather + combine. +# +# The DP TP-MoE path (deepseek_v4) gathers local hidden -> a global buffer +# before the experts and reduce-scatters back after. For TBO we run those two +# collectives on a single shared comm stream (mirroring the mori dispatcher's +# _comm_stream) and return a CUDA event, so the op engine can yield and let the +# OTHER ubatch's attn+MoE compute run on the compute stream while this ubatch's +# gather/combine proceeds on the comm stream. Both ubatches share ONE comm +# stream -> their collectives serialize in-order (no concurrent-collective +# deadlock on the RCCL communicator), each overlapping the other's compute. +# --------------------------------------------------------------------------- +_DP_TBO_COMM_STREAM: Optional[torch.cuda.Stream] = None + + +def get_dp_tbo_comm_stream() -> torch.cuda.Stream: + global _DP_TBO_COMM_STREAM + if _DP_TBO_COMM_STREAM is None: + _DP_TBO_COMM_STREAM = torch.cuda.Stream() + return _DP_TBO_COMM_STREAM + + +# Persistent reusable CUDA events for non-EP DP TBO, keyed by (kind, subbatch). +# CRITICAL: do NOT create a fresh event per gather/combine -- that is ~244 new +# torch.cuda.Event per forward (61 layers x 2 ubatches x 2), and the HSA signal +# pool is exhausted after a few hundred forwards -> HSA_STATUS_ERROR_OUT_OF_RESOURCES +# ("...create internal OS-specific events"). Reuse one event per (kind, subbatch) +# and just re-record it (mirrors the mori CommStreamPool event reuse). +_TBO_EVENT_POOL: dict = {} + + +def _tbo_event(key) -> torch.cuda.Event: + ev = _TBO_EVENT_POOL.get(key) + if ev is None: + ev = torch.cuda.Event() + _TBO_EVENT_POOL[key] = ev + return ev + + +def dp_gather_partial_async( + global_tokens: torch.Tensor, + local_tokens: torch.Tensor, + forward_batch: ForwardBatch, + event_key=("gather", 0), +) -> torch.cuda.Event: + """Launch `dp_gather_partial` (all_gatherv) on the shared DP TBO comm stream; + re-record + return a PERSISTENT event (keyed by `event_key`) that fires when + the gather completes. Caller yields, then `compute_stream.wait_event(ev)` + before reading `global_tokens`.""" + comm = get_dp_tbo_comm_stream() + compute = torch.cuda.current_stream() + # Keep buffers alive across streams (caching allocator). + local_tokens.record_stream(comm) + global_tokens.record_stream(comm) + ev = _tbo_event(event_key) + with torch.cuda.stream(comm): + comm.wait_stream(compute) # inputs were produced on the compute stream + dp_gather_partial(global_tokens, local_tokens, forward_batch) + ev.record(comm) + return ev + + +# Persistent grow-only buffers for non-EP DP TBO, keyed by (kind, tbo_subbatch). +# Reused across ALL layers (and forwards) so the caching allocator does not churn +# a fresh per-layer `torch.empty` for the 8x DP-gather / combine buffers. That +# churn (different sizes per forward x 2 ubatches x 61 layers, kept alive by the +# comm-stream record_stream) ballooned `reserved` to ~270GB and tripped +# HSA_STATUS_ERROR_OUT_OF_RESOURCES at large prefill chunks, even though the live +# (allocated) working set was only ~10GB. +_TBO_PERSIST_BUF: dict = {} + + +def get_tbo_persistent_buffer( + key, rows: int, hidden: int, dtype: torch.dtype, device +) -> torch.Tensor: + """Return a [rows, hidden] view of a grow-only persistent buffer for `key`. + Reallocates only when the request exceeds the cached capacity / changes + dtype|hidden. Caller must treat the returned view as scratch (overwritten).""" + buf = _TBO_PERSIST_BUF.get(key) + cap = 0 if buf is None else buf.shape[0] + if buf is None or rows > cap or buf.shape[1] != hidden or buf.dtype != dtype: + new_rows = max(rows, cap) + buf = torch.empty((new_rows, hidden), dtype=dtype, device=device) + _TBO_PERSIST_BUF[key] = buf + return buf[:rows] + + +def dp_reduce_scatterv_async( + output_local: torch.Tensor, + global_tokens: torch.Tensor, + sizes: List[int], + event_key=("combine", 0), +) -> torch.cuda.Event: + """Launch the variable-length reduce_scatterv (combine) on the shared DP TBO + comm stream; re-record + return a PERSISTENT event (keyed by `event_key`). + Matches the gatherv (SUM_LEN) path.""" + comm = get_dp_tbo_comm_stream() + compute = torch.cuda.current_stream() + ev = _tbo_event(event_key) + with torch.cuda.stream(comm): + comm.wait_stream(compute) + get_tp_group().reduce_scatterv(global_tokens, output=output_local, sizes=sizes) + ev.record(comm) + return ev + + def attn_tp_reduce_scatter_tensor(output: torch.Tensor, input: torch.Tensor): return get_attention_tp_group().reduce_scatter_tensor(output, input) diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 82d421a96..12b9b112e 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -1452,6 +1452,14 @@ class DeepseekV2MoE(nn.Module): router_logits = state.pop("router_logits") hidden_states = state.hidden_states_mlp_input + # Hash MoE layers (e.g. DeepSeek-V4) route on input_ids; forward_deepep + # passes them as a topk kwarg. The per-ubatch forward_batch.input_ids is + # already sliced+padded to match hidden_states rows (and equals the + # global ids under EP dp-attention). No-op for non-hash models. + topk_kwargs = {} + if getattr(self, "is_hash", False): + topk_kwargs["input_ids"] = state.forward_batch.input_ids + if router_logits is not None: with get_global_expert_distribution_recorder().with_current_layer( self.layer_id @@ -1463,6 +1471,7 @@ class DeepseekV2MoE(nn.Module): expert_location_dispatch_info=ExpertLocationDispatchInfo.init_new( layer_id=self.layer_id, ), + **topk_kwargs, ) else: state.topk_output = self.topk.empty_topk_output( diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 636ab7e14..f470bcd60 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -53,15 +53,21 @@ from sglang.srt.layers.deepseek_v4_rope import ( ) from sglang.srt.layers.dp_attention import ( _DpGatheredBufferWrapper, + _tbo_event, attn_tp_all_gather, attn_tp_all_reduce, dp_gather_partial, dp_gather_replicate, dp_reduce_scatter_tensor, + dp_reduce_scatterv_async, dp_scatter, get_dp_global_num_tokens, + get_dp_tbo_comm_stream, get_global_dp_buffer, + get_global_dp_buffer_len, get_local_dp_buffer, + get_local_dp_buffer_len, + get_tbo_persistent_buffer, is_dp_attention_enabled, is_dp_gatherv_active, ) @@ -1108,6 +1114,21 @@ class MQALayer(nn.Module): return o + # ---- TBO op decomposition (prefill two-batch-overlap) ---- + def op_attn(self, state): + """Run the attention forward as a single TBO op. + + Consumes the post-input-norm hidden states produced by + ``DeepseekV4DecoderLayer.op_mhc_prepare_attn`` and stores the attention + output for ``op_mhc_post_attn_pre_mlp``. + """ + state.hidden_states_after_attn = self.forward( + x=state.pop("hidden_states_after_input_norm"), + positions=state.positions, + forward_batch=state.forward_batch, + x_quant=state.pop("attn_x_quant"), + ) + class DeepseekV4DecoderLayer(nn.Module): def __init__( @@ -1600,6 +1621,217 @@ class DeepseekV4DecoderLayer(nn.Module): # cross-layer fusion, and the final layer is completed in DeepseekV4Model. return hidden_states, residual, post, comb + # ------------------------------------------------------------------ + # TBO op decomposition (prefill two-batch-overlap, EP / mori path) + # + # These mirror the NON-fused branch of ``forward`` (cross-layer mHC + # fusion is disabled under TBO, so every layer is self-contained), split + # into ops so the operations engine can overlap one ubatch's MoE a2a + # dispatch/combine with the other ubatch's attention + expert GEMM. + # The MoE ops themselves (op_gate / op_select_experts / op_dispatch_a/b / + # op_experts / op_combine_a/b / op_shared_experts / op_output) are reused + # as-is from ``self.mlp`` (DeepseekV2MoE) — they decompose ``forward_deepep``. + # ------------------------------------------------------------------ + def op_mhc_prepare_attn( + self, + state, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + residual: Optional[torch.Tensor] = None, + tbo_subbatch_index: Optional[int] = None, + **kwargs, + ): + # Non-fused attention-side mHC pre + input layernorm. + attn_residual = hidden_states + hidden_states, post, comb, norm_fused = self.hc_pre( + hidden_states, + self.hc_attn_fn, + self.hc_attn_scale, + self.hc_attn_base, + norm=self.input_layernorm, + forward_batch=forward_batch, + ) + if not norm_fused: + if _use_aiter and _is_gfx95_supported: + x_quant, hidden_states = _fused_rmsnorm_fp8_quant( + hidden_states, + self.input_layernorm.weight, + self.rms_norm_eps, + ) + else: + hidden_states = self.input_layernorm(hidden_states) + x_quant = None + else: + x_quant = None + + state.attn_residual = attn_residual + state.attn_post = post + state.attn_comb = comb + state.hidden_states_after_input_norm = hidden_states + state.attn_x_quant = x_quant + # mori's op_output slices final_hidden_states[:num_tokens]. + if get_moe_a2a_backend().is_mori(): + state.num_tokens = attn_residual.shape[0] + state.update( + dict( + forward_batch=forward_batch, + positions=positions, + tbo_subbatch_index=tbo_subbatch_index, + ) + ) + + def op_mhc_post_attn_pre_mlp(self, state): + # Close the attention mHC (hc_post), then open the FFN-side mHC pre + + # post-attention layernorm. Produces the 2D MoE input. + hidden_states = self.hc_post( + state.pop("hidden_states_after_attn"), + state.pop("attn_residual"), + state.pop("attn_post"), + state.pop("attn_comb"), + ) + ffn_residual = hidden_states + hidden_states, post, comb, norm_fused = self.hc_pre( + hidden_states, + self.hc_ffn_fn, + self.hc_ffn_scale, + self.hc_ffn_base, + norm=self.post_attention_layernorm, + forward_batch=state.forward_batch, + ) + if not norm_fused: + hidden_states = self.post_attention_layernorm(hidden_states) + state.ffn_residual = ffn_residual + state.ffn_post = post + state.ffn_comb = comb + state.hidden_states_mlp_input = hidden_states + + def op_mhc_postprocess(self, state): + # Close the FFN mHC (hc_post) and emit the next layer's input dict. + hidden_states = self.hc_post( + state.pop("hidden_states_mlp_output"), + state.pop("ffn_residual"), + state.pop("ffn_post"), + state.pop("ffn_comb"), + ) + output = dict( + positions=state.positions, + hidden_states=hidden_states, + # DSV4 non-fused layers carry no residual across layers; the key is + # required by the next layer's op_mhc_prepare_attn (ignored) and by + # _model_forward_tbo_merge_outputs (None -> None). + residual=None, + forward_batch=state.forward_batch, + tbo_subbatch_index=state.tbo_subbatch_index, + ) + state.clear( + expect_keys={ + "positions", + "forward_batch", + "tbo_subbatch_index", + } + ) + return output + + # ------------------------------------------------------------------ + # Non-EP (DP TP-MoE) TBO ops. Overlap the DP all_gatherv (pre-MoE gather) + # + reduce_scatterv (post-MoE combine) with the OTHER ubatch's attn+MoE + # compute. Used when moe_a2a_backend is "none" (DP-attention, TP-MoE) — + # the path ATOM uses for DSV4 (+~7.7% prefill). Replaces the EP mori + # op_dispatch/op_combine. op_mhc_* and op_attn are reused (local hidden). + # ------------------------------------------------------------------ + def op_gather_a(self, state): + # Launch the all_gatherv (local hidden -> global buffer) + the input_ids + # replicate-gather on the shared comm stream; record an event. + fb = state.forward_batch + local = state.pop("hidden_states_mlp_input") # LOCAL [M_local, hidden] + # Shared-expert-local: compute on LOCAL hidden before the gather; added + # back after the combine (same as the non-fused forward). Skipped in the + # global MoE via skip_shared_experts. + do_shared_local = ( + _SHARED_EXPERT_LOCAL + and getattr(self.mlp, "shared_experts", None) is not None + and getattr(self.mlp, "_shared_expert_tp1", False) + ) + state.do_shared_local = do_shared_local + state.shared_local = ( + self.mlp._forward_shared_experts(local) + if (do_shared_local and local.shape[0] > 0) + else None + ) + # Persistent grow-only scratch (keyed per ubatch) instead of a fresh + # torch.empty each layer -> stops the allocator's `reserved` from + # ballooning at large prefill chunks. input_ids_global is gathered ONCE + # per ubatch in _forward_layers_tbo (cached on fb), not here. + sub = state.tbo_subbatch_index + global_rows = get_global_dp_buffer_len() + global_hidden = get_tbo_persistent_buffer( + ("gh", sub), global_rows, local.shape[1], local.dtype, local.device + ) + comm = get_dp_tbo_comm_stream() + compute = torch.cuda.current_stream() + with torch.cuda.stream(comm): + comm.wait_stream(compute) + dp_gather_partial(global_hidden, local, fb) + state.gather_event = _tbo_event(("gather", sub)) + state.gather_event.record(comm) + state.gather_keepalive = local + state.global_hidden = global_hidden + + def op_gather_b(self, state): + torch.cuda.current_stream().wait_event(state.pop("gather_event")) + # Compute now ordered after the gather -> the gather input is safe to + # release (freed on the compute stream, no record_stream deferral). + state.pop("gather_keepalive") + + def op_moe(self, state): + # MoE (gate/topk/experts) on the GLOBAL gathered buffer. use_reduce_scatter + # skips the MoE-internal all_reduce (we reduce_scatterv in op_combine). + fb = state.forward_batch + global_hidden = state.pop("global_hidden") + global_ids = fb._tbo_global_input_ids + state.global_expert_out = self.mlp( + global_hidden, + fb, + use_reduce_scatter=True, + input_ids=global_ids, + input_ids_global=global_ids, + skip_shared_experts=state.do_shared_local, + ) + + def op_combine_a(self, state): + # Launch reduce_scatterv (global partial expert sums -> per-rank local) on + # the comm stream; record an event. Symmetric inverse of the all_gatherv. + global_out = state.pop("global_expert_out") + local_out = get_tbo_persistent_buffer( + ("lo", state.tbo_subbatch_index), + get_local_dp_buffer_len(), + global_out.shape[1], + global_out.dtype, + global_out.device, + ) + state.combine_event = dp_reduce_scatterv_async( + local_out, + global_out, + get_dp_global_num_tokens(), + event_key=("combine", state.tbo_subbatch_index), + ) + state.local_out = local_out + # Keep the (variable-size) MoE output alive until op_combine_b waits on + # the combine event (replaces record_stream; avoids reserved churn). + state.combine_keepalive = global_out + + def op_combine_b(self, state): + torch.cuda.current_stream().wait_event(state.pop("combine_event")) + state.pop("combine_keepalive") + hidden = state.pop("local_out") + shared_local = state.pop("shared_local") + state.pop("do_shared_local") + if shared_local is not None: + n = hidden.shape[0] + hidden = hidden + shared_local[:n] + state.hidden_states_mlp_output = hidden + class DeepseekV4Model(nn.Module): fall_back_to_pt_during_load = False @@ -1695,6 +1927,116 @@ class DeepseekV4Model(nn.Module): y = torch.sum(pre.unsqueeze(-1) * x.view(shape), dim=1) return y.to(dtype) + def _can_run_tbo(self, forward_batch: ForwardBatch) -> bool: + """DSV4 prefill-only two-batch-overlap gate. + + TBO batch prep (tbo_split_seq_index / tbo_children) is populated + model-agnostically when --enable-two-batch-overlap is set and the + DP-attention preparer allows it (mori `normal` mode permits prefill + TBO). We additionally restrict to: prefill (EXTEND), single PP, and the + non-CP path, which is the only case the DSV4 op strategy implements. + """ + from sglang.srt.layers.moe import is_tbo_enabled + + return ( + is_tbo_enabled() + and forward_batch.can_run_tbo + and forward_batch.tbo_children is not None + and forward_batch.global_forward_mode is not None + and forward_batch.global_forward_mode.is_extend() + and not dsa_use_prefill_cp(forward_batch) + and self.pp_group.world_size == 1 + ) + + def _forward_layers_tbo( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + ) -> torch.Tensor: + from sglang.srt.batch_overlap.operations import execute_overlapped_operations + from sglang.srt.batch_overlap.operations_strategy import OperationsStrategy + from sglang.srt.batch_overlap.two_batch_overlap import ( + _model_forward_filter_inputs, + _model_forward_tbo_merge_outputs, + ) + + layers = [self.layers[i] for i in range(self.start_layer, self.end_layer)] + operations_strategy = OperationsStrategy.init_new_tbo( + layers, forward_batch.global_forward_mode + ) + + # Split the per-rank batch into the 2 ubatches (token-range slice + pad + # to tbo_padded_len). residual is unused by the DSV4 non-fused layer ops. + inputs_arr = [ + _model_forward_filter_inputs( + hidden_states=hidden_states, + residual=None, + positions=positions, + output_forward_batch=child, + tbo_subbatch_index=idx, + ) + for idx, child in enumerate(forward_batch.tbo_children) + ] + + # Non-EP DP TP-MoE: the per-ubatch DP gather/combine (op_gather/op_combine) + # needs each ubatch's per-rank token counts, but tbo_padded_len is computed + # per-rank locally (not synced). All-gather both ubatches' padded lengths + # once across DP ranks, then populate each child's global_num_tokens + + # global_dp_buffer_len so the gatherv/reduce_scatterv buffers size correctly. + if get_moe_a2a_backend().is_none() and get_parallel().attn_dp_size > 1: + tp_group = get_tp_group() + world = tp_group.world_size + children = forward_batch.tbo_children + local_lens = torch.tensor( + [int(c.tbo_padded_len) for c in children], + dtype=torch.int64, + device=hidden_states.device, + ) + gathered = torch.empty( + (world, local_lens.shape[0]), + dtype=torch.int64, + device=hidden_states.device, + ) + tp_group.all_gather_into_tensor(gathered, local_lens) + gathered_cpu = gathered.tolist() + rank = tp_group.rank_in_group + for idx, child in enumerate(children): + sizes = [gathered_cpu[r][idx] for r in range(world)] + child.global_num_tokens_cpu = sizes + child.global_num_tokens_gpu = gathered[:, idx].contiguous() + child.global_dp_buffer_len = sum(sizes) + # Gather the ubatch's input_ids -> global ONCE here (cached on the + # child) instead of per-layer in op_gather_a. The hash MoE reads + # the SAME global ids every layer, so 61x2 per-layer all_gatherv of + # VARYING size (-> RCCL registers a new internal buffer per size -> + # HSA_STATUS_ERROR_OUT_OF_RESOURCES) collapses to 1 per ubatch. + local_ids = child.input_ids + rows = sizes[rank] + if local_ids.shape[0] < rows: + padded_ids = local_ids.new_zeros((rows,)) + padded_ids[: local_ids.shape[0]] = local_ids + elif local_ids.shape[0] > rows: + padded_ids = local_ids[:rows] + else: + padded_ids = local_ids + gids = torch.empty( + (sum(sizes),), dtype=local_ids.dtype, device=local_ids.device + ) + tp_group.all_gatherv(padded_ids, sizes=sizes, output=gids) + child._tbo_global_input_ids = gids + + outputs_arr = execute_overlapped_operations( + inputs_arr=inputs_arr, + operations_arr=[operations_strategy.operations] * 2, + delta_stages=[0, operations_strategy.tbo_delta_stages], + ) + + hidden_states, _ = _model_forward_tbo_merge_outputs( + outputs_arr[0], outputs_arr[1], hidden_states.shape[0] + ) + return hidden_states + def forward( self, input_ids: torch.Tensor, @@ -1740,32 +2082,41 @@ class DeepseekV4Model(nn.Module): if hasattr(forward_batch, _attr): delattr(forward_batch, _attr) - use_fused = self.use_fused_mhc_post_pre - prev_residual, prev_post, prev_comb = None, None, None - last_layer = None - for i in range(self.start_layer, self.end_layer): - layer = self.layers[i] - last_layer = layer - ctx = ( - nullcontext() - if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - else get_global_expert_distribution_recorder().with_current_layer(i) + if self._can_run_tbo(forward_batch): + # Two-batch-overlap prefill (EP / mori). Cross-layer mHC fusion is + # disabled here (each layer self-contained), so no trailing hc_post. + hidden_states = self._forward_layers_tbo( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, ) - with ctx: - hidden_states, prev_residual, prev_post, prev_comb = layer( - positions=positions, - hidden_states=hidden_states, - forward_batch=forward_batch, - input_ids=input_ids, - input_ids_global=input_ids_global, - prev_residual=prev_residual, - prev_post=prev_post, - prev_comb=prev_comb, + else: + use_fused = self.use_fused_mhc_post_pre + prev_residual, prev_post, prev_comb = None, None, None + last_layer = None + for i in range(self.start_layer, self.end_layer): + layer = self.layers[i] + last_layer = layer + ctx = ( + nullcontext() + if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) + else get_global_expert_distribution_recorder().with_current_layer(i) + ) + with ctx: + hidden_states, prev_residual, prev_post, prev_comb = layer( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, + input_ids=input_ids, + input_ids_global=input_ids_global, + prev_residual=prev_residual, + prev_post=prev_post, + prev_comb=prev_comb, + ) + if use_fused and last_layer is not None: + hidden_states = last_layer.hc_post( + hidden_states, prev_residual, prev_post, prev_comb ) - if use_fused and last_layer is not None: - hidden_states = last_layer.hc_post( - hidden_states, prev_residual, prev_post, prev_comb - ) # CP all-gather only on the last PP rank; PP IPC carries CP-split tensors. if self.pp_group.is_last_rank and dsa_use_prefill_cp(forward_batch): diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index ae6122884..5492cee7a 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -6434,6 +6434,22 @@ class ServerArgs: self._mamba_cache_chunk_size = max(chunk_size, self.page_size) return self._mamba_cache_chunk_size + def _check_two_batch_overlap(self): + # With no EP a2a backend, two-batch-overlap is only valid on the non-EP + # DP TP-MoE path (overlapping the DP all_gatherv / reduce_scatterv with + # the other ubatch's compute), which requires DP attention. Enabling it + # there needs no extra opt-in env flag. + if ( + self.enable_two_batch_overlap + and self.moe_a2a_backend == "none" + and not self.enable_dp_attention + ): + raise ValueError( + "When enabling two batch overlap without an EP a2a backend " + "(moe_a2a_backend='none'), --enable-dp-attention is required " + "(DeepSeek-V4 non-EP DP TBO path)." + ) + def check_server_args(self): # Check parallel size constraints assert ( @@ -6581,11 +6597,8 @@ class ServerArgs: "--export-metrics-to-file-dir is required when --export-metrics-to-file is enabled" ) - # Check two batch overlap - if self.enable_two_batch_overlap and self.moe_a2a_backend == "none": - raise ValueError( - "When enabling two batch overlap, moe_a2a_backend cannot be 'none'." - ) + # Check two batch overlap backend requirement. + self._check_two_batch_overlap() # Check communications compression if self.enable_quant_communications and self.tp_size == 1: diff --git a/test/registered/amd/test_deepseek_v4_flash_fp8_tbo.py b/test/registered/amd/test_deepseek_v4_flash_fp8_tbo.py new file mode 100644 index 000000000..01f2ab613 --- /dev/null +++ b/test/registered/amd/test_deepseek_v4_flash_fp8_tbo.py @@ -0,0 +1,163 @@ +"""MI35x DeepSeek-V4-Flash FP8 + non-EP DP two-batch-overlap (TBO) test (8-GPU) + +End-to-end accuracy test for DeepSeek-V4-Flash (285B) FP8 with the non-EP DP +two-batch-overlap path on MI35x ROCm 7.2. + +TBO here is the DP-attention TP-MoE variant (moe_a2a_backend='none'): it overlaps +one micro-batch's DP all_gatherv (pre-MoE gather) + reduce_scatterv (post-MoE +combine) with the other micro-batch's attention + expert compute (prefill only). +Enabled purely via `--enable-dp-attention` + `--enable-two-batch-overlap` (no opt-in +env). This test guards that TBO does not regress GSM8K accuracy and that the DP TBO +server launches + runs to completion (exercises op_gather/op_moe/op_combine and the +event+ref combine-buffer lifetime that fixed the reserved-memory OOM at mem0.9). + +Unlike the CPU-only server-args guard unit test (TestTwoBatchOverlapBackend), this +actually runs the TBO forward on the real model — which only DeepSeek-V4 implements, +so it needs the real 8-GPU model (a dummy model path would not exercise TBO). + +Registry: nightly-amd-8-gpu-mi35x-deepseek-v4-flash suite +""" + +import os +import unittest +from types import SimpleNamespace + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_amd_ci +from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k +from sglang.test.test_utils import ( + DEFAULT_URL_FOR_TEST, + CustomTestCase, + is_in_ci, + popen_launch_server, + write_github_step_summary, +) + +register_amd_ci( + est_time=7200, suite="nightly-amd-8-gpu-mi35x-deepseek-v4-flash", nightly=True +) + +DEEPSEEK_V4_FLASH_FP8_MODEL_PATH = os.environ.get( + "DEEPSEEK_V4_FP8_MODEL_PATH", "sgl-project/DeepSeek-V4-Flash-FP8" +) +SERVER_LAUNCH_TIMEOUT = 3600 +FLASHMLA_BACKEND = os.environ.get("SGLANG_HACK_FLASHMLA_BACKEND", "unified_kv_triton") + +# DSV4 fused-kernel optimal set (mirrors the validated dp-tbo launch config). +# The DP + TBO forward path is sensitive to these; the non-TBO tp8 test can use a +# leaner set, but DP TBO needs the full DSV4 opt env or the MoE/attn kernels hit +# shape mismatches at warmup. +COMMON_ENV_VARS = { + "SGLANG_DEFAULT_THINKING": "1", + "SGLANG_DSV4_REASONING_EFFORT": "max", + "SGLANG_OPT_DEEPGEMM_HC_PRENORM": "false", + "SGLANG_USE_AITER": "1", + "SGLANG_USE_ROCM700A": "0", + "SGLANG_OPT_USE_FUSED_COMPRESS": "true", + "SGLANG_HACK_FLASHMLA_BACKEND": FLASHMLA_BACKEND, + "SGLANG_OPT_FP8_WO_A_GEMM": "false", + "SGLANG_OPT_USE_JIT_INDEXER_METADATA": "false", + "SGLANG_OPT_USE_TOPK_V2": "false", + "SGLANG_OPT_USE_AITER_INDEXER": "true", + "SGLANG_OPT_USE_TILELANG_INDEXER": "false", + "SGLANG_OPT_USE_TILELANG_MHC_PRE": "false", + "SGLANG_OPT_USE_TILELANG_MHC_POST": "false", + "SGLANG_FP8_PAGED_MQA_LOGITS_TORCH": "1", + "SGLANG_OPT_USE_FUSED_COMPRESS_TRITON": "true", + "SGLANG_OPT_USE_MULTI_STREAM_OVERLAP": "false", + "SGLANG_ROCM_USE_MULTI_STREAM": "false", + "AITER_BF16_FP8_MOE_BOUND": "0", + "SGLANG_EAGER_INPUT_NO_COPY": "true", + # DP TP-MoE collective path that non-EP DP TBO overlaps. + "SGLANG_DP_USE_GATHERV": "1", + "SGLANG_DP_USE_REDUCE_SCATTER": "1", + "SGLANG_SHARED_EXPERT_TP1": "1", + "SGLANG_DP_SHARED_EXPERT_LOCAL": "1", + # ROCm HSA-resource stability for TBO at high concurrency. + "GPU_MAX_HW_QUEUES": "5", + # FP8 variant + "SGLANG_DSV4_FP4_EXPERTS": "false", +} + + +class TestDeepseekV4FlashFp8Tbo(CustomTestCase): + @classmethod + def setUpClass(cls): + cls.model = DEEPSEEK_V4_FLASH_FP8_MODEL_PATH + cls.base_url = DEFAULT_URL_FOR_TEST + + env = os.environ.copy() + env.update(COMMON_ENV_VARS) + + other_args = [ + "--trust-remote-code", + "--tp", + "8", + # DP attention + TBO: non-EP DP TP-MoE two-batch-overlap. DP TBO is + # selected because moe_a2a_backend stays 'none'; no opt-in env needed. + "--dp", + "8", + "--enable-dp-attention", + "--enable-prefill-delayer", + "--enable-two-batch-overlap", + "--disable-radix-cache", + "--attention-backend", + "dsv4", + "--kv-cache-dtype", + "fp8_e4m3", + "--max-running-requests", + "512", + "--cuda-graph-max-bs", + "512", + "--page-size", + "256", + "--mem-fraction-static", + "0.90", + "--swa-full-tokens-ratio", + "0.15", + # global chunk; DP-attention divides by dp_size=8 -> 8192/rank. + "--chunked-prefill-size", + "65536", + "--disable-shared-experts-fusion", + "--tool-call-parser", + "deepseekv4", + "--reasoning-parser", + "deepseek-v4", + ] + + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=SERVER_LAUNCH_TIMEOUT, + other_args=other_args, + env=env, + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_gsm8k_tbo(self): + args = SimpleNamespace( + num_shots=8, + data_path=None, + num_questions=1319, + parallel=1319, + max_new_tokens=512, + host="http://127.0.0.1", + port=int(self.base_url.split(":")[-1]), + ) + metrics = run_eval_few_shot_gsm8k(args) + print(f"{metrics=}") + + if is_in_ci(): + write_github_step_summary( + f"### test_gsm8k_tbo (deepseek-v4-flash-fp8 DP+TBO, {FLASHMLA_BACKEND})\n" + f'{metrics["accuracy"]=:.3f}\n' + ) + # TBO must not regress accuracy vs the non-TBO baseline (>0.91). + self.assertGreater(metrics["accuracy"], 0.91) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/amd/test_deepseek_v4_pro_fp4_tbo.py b/test/registered/amd/test_deepseek_v4_pro_fp4_tbo.py new file mode 100644 index 000000000..65995474e --- /dev/null +++ b/test/registered/amd/test_deepseek_v4_pro_fp4_tbo.py @@ -0,0 +1,151 @@ +"""MI35x DeepSeek-V4-Pro FP4 + non-EP DP two-batch-overlap (TBO) test (8-GPU) + +End-to-end accuracy test for DeepSeek-V4-Pro (1.6T) FP4 with the non-EP DP +two-batch-overlap path on MI35x ROCm 7.2. + +TBO here is the DP-attention TP-MoE variant (moe_a2a_backend='none'): it overlaps +one micro-batch's DP all_gatherv (pre-MoE gather) + reduce_scatterv (post-MoE +combine) with the other micro-batch's attention + expert compute (prefill only). +Enabled purely via `--enable-dp-attention` + `--enable-two-batch-overlap` (no opt-in +env). This test guards that TBO does not regress GSM8K accuracy and that the DP TBO +server launches + runs (exercises op_gather/op_moe/op_combine and the event+ref +combine-buffer lifetime that fixed the reserved-memory OOM at mem0.9). + +Unlike the CPU-only server-args guard unit test (TestTwoBatchOverlapBackend), this +runs the TBO forward on the real model — which only DeepSeek-V4 implements — so it +needs the real 8-GPU model (a dummy model path would not exercise TBO). Uses the FP4 +Pro model (fp4 routed experts); do NOT force SGLANG_DSV4_FP4_EXPERTS=false here or +the expert weights are read at the wrong (fp8) shape. + +Registry: nightly-amd-8-gpu-mi35x-deepseek-v4-pro suite +""" + +import os +import unittest +from types import SimpleNamespace + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_amd_ci +from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k +from sglang.test.test_utils import ( + DEFAULT_URL_FOR_TEST, + CustomTestCase, + is_in_ci, + popen_launch_server, + write_github_step_summary, +) + +register_amd_ci( + est_time=14400, suite="nightly-amd-8-gpu-mi35x-deepseek-v4-pro", nightly=True +) + +DEEPSEEK_V4_PRO_FP4_MODEL_PATH = os.environ.get( + "DEEPSEEK_V4_PRO_MODEL_PATH_FP4", "deepseek-ai/DeepSeek-V4-Pro" +) +# Pro is 1.6T; weight load + warmup is much longer than Flash 285B. +SERVER_LAUNCH_TIMEOUT = 5400 +FLASHMLA_BACKEND = os.environ.get("SGLANG_HACK_FLASHMLA_BACKEND", "unified_kv_triton") + +COMMON_ENV_VARS = { + "SGLANG_DEFAULT_THINKING": "1", + "SGLANG_DSV4_REASONING_EFFORT": "max", + "SGLANG_USE_ROCM700A": "0", + "SGLANG_HACK_FLASHMLA_BACKEND": FLASHMLA_BACKEND, + "AITER_BF16_FP8_MOE_BOUND": "0", + # DP TP-MoE collective path that non-EP DP TBO overlaps. + "SGLANG_DP_USE_GATHERV": "1", + "SGLANG_DP_USE_REDUCE_SCATTER": "1", + "SGLANG_SHARED_EXPERT_TP1": "1", + "SGLANG_DP_SHARED_EXPERT_LOCAL": "1", + # ROCm HSA-resource stability for TBO at high concurrency. + "GPU_MAX_HW_QUEUES": "5", +} + +# FP4 variant +FP4_ENV_VARS = { + "SGLANG_DSV4_FP4_EXPERTS": "true", +} + + +class TestDeepseekV4ProFp4Tbo(CustomTestCase): + @classmethod + def setUpClass(cls): + cls.model = DEEPSEEK_V4_PRO_FP4_MODEL_PATH + cls.base_url = DEFAULT_URL_FOR_TEST + + env = os.environ.copy() + env.update(COMMON_ENV_VARS) + env.update(FP4_ENV_VARS) + + other_args = [ + "--trust-remote-code", + "--tp", + "8", + # DP attention + TBO: non-EP DP TP-MoE two-batch-overlap. DP TBO is + # selected because moe_a2a_backend stays 'none'; no opt-in env needed. + "--dp", + "8", + "--enable-dp-attention", + "--enable-prefill-delayer", + "--enable-two-batch-overlap", + "--disable-radix-cache", + "--attention-backend", + "dsv4", + "--kv-cache-dtype", + "fp8_e4m3", + "--max-running-requests", + "512", + "--cuda-graph-max-bs", + "512", + "--page-size", + "256", + "--mem-fraction-static", + "0.90", + "--swa-full-tokens-ratio", + "0.15", + # global chunk; DP-attention divides by dp_size=8 -> 8192/rank. + "--chunked-prefill-size", + "65536", + "--disable-shared-experts-fusion", + "--tool-call-parser", + "deepseekv4", + "--reasoning-parser", + "deepseek-v4", + ] + + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=SERVER_LAUNCH_TIMEOUT, + other_args=other_args, + env=env, + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_gsm8k_tbo(self): + args = SimpleNamespace( + num_shots=8, + data_path=None, + num_questions=1319, + parallel=1319, + max_new_tokens=512, + host="http://127.0.0.1", + port=int(self.base_url.split(":")[-1]), + ) + metrics = run_eval_few_shot_gsm8k(args) + print(f"{metrics=}") + + if is_in_ci(): + write_github_step_summary( + f"### test_gsm8k_tbo (deepseek-v4-pro-fp4 DP+TBO, {FLASHMLA_BACKEND})\n" + f'{metrics["accuracy"]=:.3f}\n' + ) + # TBO must not regress accuracy vs the non-TBO baseline (>0.91). + self.assertGreater(metrics["accuracy"], 0.91) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 39e46700f..ed8ac52b4 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -1364,5 +1364,48 @@ class TestSamplingBackendTokenOracleEnvGate(CustomTestCase): self.assertEqual(parsed.sampling_backend, "token_oracle") +class TestTwoBatchOverlapBackend(CustomTestCase): + """Non-EP DP two-batch-overlap backend requirement. + + With no EP a2a backend (moe_a2a_backend='none'), --enable-two-batch-overlap + is only valid on the DeepSeek-V4 non-EP DP TP-MoE path (overlapping the DP + all_gatherv / reduce_scatterv with the other ubatch's compute), which + requires --enable-dp-attention. This replaced the removed opt-in + SGLANG_ENABLE_DP_TBO env: enabling DP TBO now needs no extra flag. + + dummy-model short-circuits __post_init__, so the guard handler is invoked + directly (same pattern as TestDeepEPWaterfillArgs).""" + + def _args(self, **overrides): + args = ServerArgs(model_path="dummy") + args.enable_two_batch_overlap = True + args.moe_a2a_backend = "none" + args.enable_dp_attention = False + for key, value in overrides.items(): + setattr(args, key, value) + return args + + def test_no_a2a_without_dp_attention_raises(self): + args = self._args(enable_dp_attention=False) + with self.assertRaisesRegex(ValueError, "enable-dp-attention"): + args._check_two_batch_overlap() + + def test_no_a2a_with_dp_attention_ok(self): + # DP TBO path is valid: --enable-dp-attention + --enable-two-batch-overlap + # with a2a backend 'none' must NOT raise (no SGLANG_ENABLE_DP_TBO needed). + args = self._args(enable_dp_attention=True) + args._check_two_batch_overlap() + + def test_ep_a2a_backend_ok_without_dp_attention(self): + # EP a2a path (e.g. deepep) overlaps dispatch/combine; the guard does not + # require dp-attention there. + args = self._args(moe_a2a_backend="deepep", enable_dp_attention=False) + args._check_two_batch_overlap() + + def test_tbo_disabled_is_noop(self): + args = self._args(enable_two_batch_overlap=False, enable_dp_attention=False) + args._check_two_batch_overlap() + + if __name__ == "__main__": unittest.main()