From 44c786679fa0bdd9917df58391705a35250c6ae3 Mon Sep 17 00:00:00 2001 From: Ming Yang Date: Fri, 4 Sep 2026 02:11:03 -0700 Subject: [PATCH] [sp] Make attention-TP sequence sharding a per-forward batch property (#37546) Co-authored-by: Claude Co-authored-by: Lianmin Zheng --- .../srt/batch_overlap/two_batch_overlap.py | 12 +- .../npu/attention/ascend_backend.py | 49 ++++-- .../npu/attention/ascend_gdn_backend.py | 10 +- .../srt/kv_canary/perturb/slot_picker.py | 6 +- .../srt/layers/attention/hpc_ops_backend.py | 2 +- python/sglang/srt/layers/logits_processor.py | 4 + python/sglang/srt/layers/radix_attention.py | 6 +- .../srt/layers/radix_linear_attention.py | 6 +- .../srt/model_executor/cpu_graph_runner.py | 4 +- .../cuda_graph_buffer_registry.py | 61 +++++--- .../srt/model_executor/forward_batch_info.py | 148 ++++++++++++------ .../sglang/srt/model_executor/model_runner.py | 36 +++-- .../runner/decode_cuda_graph_runner.py | 9 +- .../runner/prefill_cuda_graph_runner.py | 33 +++- .../model_executor/runner_utils/buffers.py | 1 + python/sglang/srt/models/deepseek_v4.py | 2 +- python/sglang/srt/models/inkling.py | 4 +- python/sglang/srt/models/nemotron_h_utils.py | 7 +- .../dspark_components/dspark_draft.py | 10 +- .../srt/speculative/frozen_kv_mtp_utils.py | 8 +- ...er_eagle_draft_extend_cuda_graph_runner.py | 2 +- .../attention_methods/dsa_attention.py | 4 +- .../attention_methods/lightning_attention.py | 2 +- .../runner_modes/split_op_runner.py | 4 +- python/sglang/test/kv_canary/fixtures.py | 4 +- .../attention/unittests/KNOWN_FAILURES.md | 4 +- .../attention/unittests/dsa/README.md | 2 +- .../attention/unittests/dsa/test_dsa.py | 2 +- .../attention/unittests/mamba/README.md | 2 +- .../attention/unittests/mamba/test_mamba2.py | 4 +- .../breakable/test_breakable_cuda_graph.py | 2 +- .../kv_canary/test_self_unit_perturb.py | 4 +- .../test_self_unit_runner_per_forward.py | 8 +- .../spec/dspark/test_dspark_dp_tier.py | 8 +- .../test_tbo_children_dummy_token_mask.py | 26 +-- .../unit/layers/test_radix_attention.py | 2 +- .../layers/test_radix_linear_attention.py | 8 +- .../test_cuda_graph_buffer_registry.py | 110 ++++++++++--- .../test_num_token_non_padded_localization.py | 70 +++++++++ .../test_prefill_cuda_graph_runner.py | 15 ++ 40 files changed, 498 insertions(+), 203 deletions(-) create mode 100644 test/registered/unit/model_executor/test_num_token_non_padded_localization.py diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index 6e9a7ba10..9a85ffe9b 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -516,7 +516,7 @@ class TboForwardBatchPreparer: batch, tbo_children_num_token_non_padded=tbo_children_num_token_non_padded, # Eager split: the children can carry a CPU count too, so the - # attention 0-token skip (which reads num_token_non_padded_cpu) + # attention 0-token skip (which reads global_num_token_non_padded_cpu) # survives the split. The cuda-graph plugin path below leaves this # None because its device buffer is refreshed per replay. tbo_children_num_token_non_padded_cpu=cls._split_num_token_non_padded( @@ -806,7 +806,11 @@ class TboForwardBatchPreparer: extend_num_tokens=extend_num_tokens, num_token_non_padded=out_num_token_non_padded, # TODO: handle it when we need TBO + DeepSeek V3.2 - num_token_non_padded_cpu=out_num_token_non_padded_cpu, + global_num_token_non_padded=None, + global_num_token_non_padded_cpu=out_num_token_non_padded_cpu, + # The child runs the same forward, so it keeps the parent's + # sharding verdict; its counts above are already per-child. + attn_tp_sequence_sharded=batch.attn_tp_sequence_sharded, tbo_split_seq_index=None, tbo_parent_token_range=(start_token_index, end_token_index), tbo_children=None, @@ -863,8 +867,8 @@ class TboForwardBatchPreparer: @staticmethod def _get_num_token_non_padded_cpu(batch: ForwardBatch) -> int: num_token_non_padded = ( - batch.num_token_non_padded_cpu - if batch.num_token_non_padded_cpu is not None + batch.global_num_token_non_padded_cpu + if batch.global_num_token_non_padded_cpu is not None else len(batch.input_ids) ) return num_token_non_padded diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py index ced598237..e8d5e1313 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py @@ -1274,9 +1274,12 @@ class AscendAttnBackend(AttentionBackend): if self.use_fia: if self._can_use_tnd(layer): num_token_padding = q.shape[0] - if num_token_padding > forward_batch.num_token_non_padded_cpu: + if ( + num_token_padding + > forward_batch.global_num_token_non_padded_cpu + ): q, k, v = [ - data[: forward_batch.num_token_non_padded_cpu] + data[: forward_batch.global_num_token_non_padded_cpu] for data in [q, k, v] ] q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim) @@ -1318,7 +1321,10 @@ class AscendAttnBackend(AttentionBackend): attn_out = attn_out.view( -1, layer.tp_q_head_num * layer.v_head_dim ) - if num_token_padding != forward_batch.num_token_non_padded_cpu: + if ( + num_token_padding + != forward_batch.global_num_token_non_padded_cpu + ): attn_out = torch.cat( [ attn_out, @@ -1411,9 +1417,12 @@ class AscendAttnBackend(AttentionBackend): """FIA supports multi-bs in the current version of CANN""" q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim) num_token_padding = q.shape[0] - if num_token_padding > forward_batch.num_token_non_padded_cpu: + if ( + num_token_padding + > forward_batch.global_num_token_non_padded_cpu + ): q, k, v = [ - data[: forward_batch.num_token_non_padded_cpu] + data[: forward_batch.global_num_token_non_padded_cpu] for data in [q, k, v] ] attn_output, _ = torch_npu.npu_fused_infer_attention_score( @@ -1439,7 +1448,10 @@ class AscendAttnBackend(AttentionBackend): -1, layer.tp_q_head_num * layer.v_head_dim ) - if num_token_padding != forward_batch.num_token_non_padded_cpu: + if ( + num_token_padding + != forward_batch.global_num_token_non_padded_cpu + ): attn_output = torch.cat( [ attn_output, @@ -1699,7 +1711,8 @@ class AscendAttnBackend(AttentionBackend): else: num_token_padding = q.shape[0] q, k, v = [ - data[: forward_batch.num_token_non_padded_cpu] for data in [q, k, v] + data[: forward_batch.global_num_token_non_padded_cpu] + for data in [q, k, v] ] q_nope, q_rope = q.split( [layer.v_head_dim, self.qk_rope_head_dim], dim=-1 @@ -1789,7 +1802,7 @@ class AscendAttnBackend(AttentionBackend): attn_output = attn_output.reshape( [-1, layer.tp_q_head_num, layer.v_head_dim] ) - if num_token_padding != forward_batch.num_token_non_padded_cpu: + if num_token_padding != forward_batch.global_num_token_non_padded_cpu: attn_output = torch.cat( [ attn_output, @@ -1862,7 +1875,8 @@ class AscendAttnBackend(AttentionBackend): else: num_token_padding = q.shape[0] q, k, v = [ - data[: forward_batch.num_token_non_padded_cpu] for data in [q, k, v] + data[: forward_batch.global_num_token_non_padded_cpu] + for data in [q, k, v] ] q_nope, q_rope = q.split( @@ -1891,7 +1905,7 @@ class AscendAttnBackend(AttentionBackend): attn_output = attn_output.reshape( -1, layer.tp_q_head_num, layer.v_head_dim ) - if num_token_padding != forward_batch.num_token_non_padded_cpu: + if num_token_padding != forward_batch.global_num_token_non_padded_cpu: attn_output = torch.cat( [ attn_output, @@ -2011,7 +2025,7 @@ class AscendAttnBackend(AttentionBackend): if not self.graph_mode: num_token_padding = query.shape[0] - query = query[: forward_batch.num_token_non_padded_cpu] + query = query[: forward_batch.global_num_token_non_padded_cpu] if self.forward_metadata.seq_lens_cpu_int is None: actual_seq_lengths_kv = self.forward_metadata.seq_lens_cpu_list @@ -2097,14 +2111,15 @@ class AscendAttnBackend(AttentionBackend): attn_output = attn_output.view(-1, layer.tp_q_head_num * layer.v_head_dim) if ( not self.graph_mode - and forward_batch.num_token_non_padded_cpu is not None - and forward_batch.num_token_non_padded_cpu != num_token_padding + and forward_batch.global_num_token_non_padded_cpu is not None + and forward_batch.global_num_token_non_padded_cpu != num_token_padding ): attn_output = torch.cat( [ attn_output, attn_output.new_zeros( - num_token_padding - forward_batch.num_token_non_padded_cpu, + num_token_padding + - forward_batch.global_num_token_non_padded_cpu, *attn_output.shape[1:], ), ], @@ -2132,8 +2147,8 @@ class AscendAttnBackend(AttentionBackend): q_rope = q_rope.view(-1, layer.tp_q_head_num, self.qk_rope_head_dim) if not self.graph_mode: num_token_padding = q.shape[0] - q_nope = q_nope[: forward_batch.num_token_non_padded_cpu] - q_rope = q_rope[: forward_batch.num_token_non_padded_cpu] + q_nope = q_nope[: forward_batch.global_num_token_non_padded_cpu] + q_rope = q_rope[: forward_batch.global_num_token_non_padded_cpu] if self.forward_metadata.seq_lens_cpu_int is None: actual_seq_lengths_kv = self.forward_metadata.seq_lens_cpu_list else: @@ -2235,7 +2250,7 @@ class AscendAttnBackend(AttentionBackend): attn_output = attn_output.view(-1, layer.tp_q_head_num * layer.v_head_dim) if ( not self.graph_mode - and forward_batch.num_token_non_padded_cpu != num_token_padding + and forward_batch.global_num_token_non_padded_cpu != num_token_padding ): attn_output = torch.cat( [ diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py index b601a52c2..c8928c06e 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py @@ -201,12 +201,12 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase): num_token_padding = mixed_qkv.shape[0] if ( not self.graph_mode - and forward_batch.num_token_non_padded_cpu != num_token_padding + and forward_batch.global_num_token_non_padded_cpu != num_token_padding ): - mixed_qkv = mixed_qkv[: forward_batch.num_token_non_padded_cpu] - a = a[: forward_batch.num_token_non_padded_cpu] - b = b[: forward_batch.num_token_non_padded_cpu] - seq_len = forward_batch.num_token_non_padded_cpu + mixed_qkv = mixed_qkv[: forward_batch.global_num_token_non_padded_cpu] + a = a[: forward_batch.global_num_token_non_padded_cpu] + b = b[: forward_batch.global_num_token_non_padded_cpu] + seq_len = forward_batch.global_num_token_non_padded_cpu batch_size = cache_indices.shape[0] draft_token_num = forward_batch.spec_info.draft_token_num diff --git a/python/sglang/srt/kv_canary/perturb/slot_picker.py b/python/sglang/srt/kv_canary/perturb/slot_picker.py index 383202dbe..86a6fdb6d 100644 --- a/python/sglang/srt/kv_canary/perturb/slot_picker.py +++ b/python/sglang/srt/kv_canary/perturb/slot_picker.py @@ -43,7 +43,9 @@ def collect_active_slots( if exclude_out_cache_loc: out_cache_loc = maybe_inaccurate_forward_batch.out_cache_loc if out_cache_loc is not None: - valid_num_tokens = maybe_inaccurate_forward_batch.num_token_non_padded_cpu + valid_num_tokens = ( + maybe_inaccurate_forward_batch.global_num_token_non_padded_cpu + ) if valid_num_tokens is None: valid_num_tokens = int(out_cache_loc.shape[0]) excluded = set( @@ -87,7 +89,7 @@ def pick_out_cache_loc_slot( total = int(out_cache_loc.shape[0]) if total <= 0: return None - valid_num_tokens = maybe_inaccurate_forward_batch.num_token_non_padded_cpu + valid_num_tokens = maybe_inaccurate_forward_batch.global_num_token_non_padded_cpu if valid_num_tokens is None: valid_num_tokens = total valid_num_tokens = int(valid_num_tokens) diff --git a/python/sglang/srt/layers/attention/hpc_ops_backend.py b/python/sglang/srt/layers/attention/hpc_ops_backend.py index 0c5a1382e..f5fb08c09 100644 --- a/python/sglang/srt/layers/attention/hpc_ops_backend.py +++ b/python/sglang/srt/layers/attention/hpc_ops_backend.py @@ -681,7 +681,7 @@ def hpc_ops_fp8_rope_store_kv( context = get_tc_piecewise_forward_context() forward_batch = context.forward_batch attention_layer = context.attention_layers[layer_id] - real_num_tokens = forward_batch.num_token_non_padded_cpu + real_num_tokens = forward_batch.global_num_token_non_padded_cpu backend = get_attn_backend() backend._run_fp8_rope_store_kv( diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index 77feb0dba..8df1a5c2f 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -271,6 +271,9 @@ class LogitsMetadata: # Whether this batch is prefill-only (no token generation needed) is_prefill_only: bool = False + # Carried from ForwardBatch so logits pruning can reconstruct the SP gather. + attn_tp_sequence_sharded: bool = False + mm_input_embeds: Optional[torch.Tensor] = None # DRAFT_EXTEND_V2: when set, lm_head and LAST hidden capture use only these @@ -324,6 +327,7 @@ class LogitsMetadata: token_ids_logprobs=forward_batch.token_ids_logprobs, extend_input_logprob_token_ids_gpu=forward_batch.extend_input_logprob_token_ids_gpu, is_prefill_only=forward_batch.is_prefill_only, + attn_tp_sequence_sharded=forward_batch.attn_tp_sequence_sharded, global_num_tokens_gpu=forward_batch.global_num_tokens_gpu, dp_local_start_pos=forward_batch.dp_local_start_pos, dp_local_num_tokens=forward_batch.dp_local_num_tokens, diff --git a/python/sglang/srt/layers/radix_attention.py b/python/sglang/srt/layers/radix_attention.py index 489d728a7..10ed79457 100644 --- a/python/sglang/srt/layers/radix_attention.py +++ b/python/sglang/srt/layers/radix_attention.py @@ -323,7 +323,7 @@ def _unified_attention_with_output_impl( forward_batch = context.forward_batch attention_layers = context.attention_layers attention_layer = attention_layers[layer_id] - real_query_num_tokens = forward_batch.num_token_non_padded_cpu + real_query_num_tokens = forward_batch.global_num_token_non_padded_cpu # Ordinary PCG attention pads Q/K/V to the same token bucket. Prefix MHA # instead supplies a fixed-capacity K/V chunk whose extent is independent # of the suffix queries, so its caller must preserve that separate extent. @@ -531,7 +531,7 @@ def unified_sparse_attention_with_output( context = get_tc_piecewise_forward_context() forward_batch = context.forward_batch attention_layer = context.attention_layers[layer_id] - real_num_tokens = forward_batch.num_token_non_padded_cpu + real_num_tokens = forward_batch.global_num_token_non_padded_cpu if real_num_tokens == 0: _zero_skipped_attn_outputs(attn_out, idx_out) @@ -602,7 +602,7 @@ def attention_with_output_extra_kwargs( context = get_tc_piecewise_forward_context() forward_batch = context.forward_batch attention_layer = context.attention_layers[layer_id] - real_num_tokens = forward_batch.num_token_non_padded_cpu + real_num_tokens = forward_batch.global_num_token_non_padded_cpu if real_num_tokens == 0: _zero_skipped_attn_outputs(output) diff --git a/python/sglang/srt/layers/radix_linear_attention.py b/python/sglang/srt/layers/radix_linear_attention.py index 58de3904f..01d829de0 100644 --- a/python/sglang/srt/layers/radix_linear_attention.py +++ b/python/sglang/srt/layers/radix_linear_attention.py @@ -117,7 +117,7 @@ class RadixLinearAttention(nn.Module): is_extend and not forward_batch.forward_mode.is_target_verify() ) real_num_tokens = ( - getattr(forward_batch, "num_token_non_padded_cpu", None) + getattr(forward_batch, "global_num_token_non_padded_cpu", None) if should_trim_padded_extend else None ) @@ -157,7 +157,9 @@ def _linear_attention_with_output_impl( forward_batch: ForwardBatch, ) -> None: """Run linear attention on the real prefix and initialize physical padding.""" - real_num_tokens = min(forward_batch.num_token_non_padded_cpu, mixed_qkv.shape[0]) + real_num_tokens = min( + forward_batch.global_num_token_non_padded_cpu, mixed_qkv.shape[0] + ) original_out_cache_loc = forward_batch.out_cache_loc # Keep the original ForwardBatch object and only narrow cache locations for diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index 6b1e34cbe..0a0c21004 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -940,8 +940,10 @@ class CPUGraphRunner: ) captured_forward_batch.encoder_out_cache_loc = None if enable_num_token_non_padded(): + # CPUGraphRunner asserts not require_gathered_buffer, so this path is + # never attn-TP sharded: LOCAL == GLOBAL. captured_forward_batch.num_token_non_padded.copy_( - forward_batch.num_token_non_padded + forward_batch.global_num_token_non_padded ) self.model_runner.attn_backend.init_forward_metadata(captured_forward_batch) diff --git a/python/sglang/srt/model_executor/cuda_graph_buffer_registry.py b/python/sglang/srt/model_executor/cuda_graph_buffer_registry.py index 2aad6d269..c1b8e1475 100644 --- a/python/sglang/srt/model_executor/cuda_graph_buffer_registry.py +++ b/python/sglang/srt/model_executor/cuda_graph_buffer_registry.py @@ -522,6 +522,8 @@ def build_decode_registry( require_gathered_buffer: bool = False, enable_prefill_cp: bool = False, require_mlp_tp_gather: bool = False, + # Per-bucket attn-TP sharded (SP) predicate; defaults to replicated. + attn_tp_sharded_fn: Callable[[int], bool] = lambda num_tokens: False, dp_size: int = 1, register_global_num_tokens: bool = True, share_pool: bool = True, @@ -648,15 +650,24 @@ def build_decode_registry( ) def _num_token_non_padded_post_fill(buf, fb, ctx): - # Gathered (DP) path overwrites the plain FB copy with this rank's - # local count; the non-gathered path keeps the copied value. - if require_gathered_buffer and not enable_prefill_cp: - buf.copy_( - compute_local_num_token_non_padded( - global_num_token_non_padded=fb.num_token_non_padded, - num_tokens_per_dp=ctx.padded_num_tokens, - ) + # init_new batches localize from the invariant GLOBAL scalar (a + # replicated / CP forward keeps the full count; sharded=False is a + # passthrough). The dense SBD draft and TBO sub-batches bypass + # init_new -- they leave the GLOBAL None and set the replicated LOCAL + # count directly, so carry that through. + if fb.global_num_token_non_padded is None: + buf.copy_(fb.num_token_non_padded) + return + sharded = not enable_prefill_cp and attn_tp_sharded_fn( + ctx.padded_num_tokens + ) + buf.copy_( + compute_local_num_token_non_padded( + global_num_token_non_padded=fb.global_num_token_non_padded, + num_tokens_per_dp=ctx.padded_num_tokens, + sharded=sharded, ) + ) slots.append( GraphSlot( @@ -664,6 +675,7 @@ def build_decode_registry( lambda _bs, _mt: (1,), torch.int32, axis="none", + copy_from_fb=False, post_fill=_num_token_non_padded_post_fill, ) ) @@ -810,6 +822,8 @@ def build_prefill_registry( enable_num_token_non_padded: bool = False, require_gathered_buffer: bool = False, enable_prefill_cp: bool = False, + # Per-bucket attn-TP sharded (SP) predicate; defaults to replicated. + attn_tp_sharded_fn: Callable[[int], bool] = lambda num_tokens: False, register_input_embeds: bool = True, share_pool: bool = True, source: Optional[Any] = None, @@ -904,23 +918,25 @@ def build_prefill_registry( ) def _prefill_num_token_non_padded_post_fill(buf, fb, ctx): - # The FB tensor was attn-TP-localized against the RAW length, but - # replay pads rows up to the capture bucket, moving the shard - # boundary — copying it verbatim would make the in-graph pad mask - # blank real tokens whenever raw < bucket. Recompute the local - # count against the padded bucket from the batch's un-adjusted - # global count, mirroring the decode registry's post_fill. - if require_gathered_buffer: - if not enable_prefill_cp: - buf.fill_( - compute_local_num_token_non_padded_cpu( - global_num_token_non_padded=fb.num_token_non_padded_cpu, - num_tokens_per_dp=ctx.padded_num_tokens, - ) + # LOCAL count for the PADDED bucket, derived from the invariant + # GLOBAL host int: replay pads rows up to the capture bucket, which + # moves the shard boundary, so the count must be recomputed against + # the bucket rather than copied. A replicated / CP forward keeps the + # full count (sharded=False is a passthrough). + if fb.global_num_token_non_padded_cpu is not None: + sharded = not enable_prefill_cp and attn_tp_sharded_fn( + ctx.padded_num_tokens + ) + buf.fill_( + compute_local_num_token_non_padded_cpu( + global_num_token_non_padded=fb.global_num_token_non_padded_cpu, + num_tokens_per_dp=ctx.padded_num_tokens, + sharded=sharded, ) + ) else: # Non-gathered FullCG still needs the live boundary rather - # than a stale/absent ForwardBatch tensor. + # than a stale/absent ForwardBatch value. buf.fill_(ctx.raw_num_tokens) slots.append( @@ -929,6 +945,7 @@ def build_prefill_registry( lambda _bs2, _mt: (1,), torch.int32, axis="none", + copy_from_fb=False, post_fill=_prefill_num_token_non_padded_post_fill, ) ) diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index f98d5855c..2b6bf7bb1 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -32,7 +32,7 @@ import warnings from dataclasses import dataclass from enum import IntEnum, auto from functools import total_ordering -from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Set, Tuple, Union +from typing import TYPE_CHECKING, Dict, List, Optional, Set, Tuple, Union import torch @@ -224,19 +224,6 @@ class CaptureHiddenMode(IntEnum): return self.value < other.value -# Predicate for whether a forward's sequence is sharded across the attn-TP group -# (vs. replicated on every rank). Injected at init; unset defaults to sharded. -_attn_tp_sequence_sharded_predicate: Optional[Callable[[int], bool]] = None - - -def register_attn_tp_sequence_sharded_predicate( - predicate: Callable[[int], bool], -) -> None: - """Register the predicate for whether a forward is sharded across attn-TP.""" - global _attn_tp_sequence_sharded_predicate - _attn_tp_sequence_sharded_predicate = predicate - - def get_server_return_hidden_states_mode() -> CaptureHiddenMode: features = get_exec().features mode = features.return_hidden_states_mode @@ -257,15 +244,14 @@ def get_required_capture_hidden_mode( return max(capture_hidden_mode, spec_capture_hidden_mode) -def _attn_tp_local_shard_bounds(num_tokens_per_dp: int) -> Tuple[int, int]: +def _attn_tp_local_shard_bounds( + num_tokens_per_dp: int, *, sharded: bool +) -> Tuple[int, int]: """(tokens_per_rank, rank_offset) of this attn-TP rank's slice of the sequence. - A replicated (non-sharded) forward puts the whole sequence on every rank, so - the slice is the full range with no offset; localizing it as a shard would - drop real tokens on non-zero ranks. + A replicated (non-sharded) forward keeps the full range on every rank. """ - predicate = _attn_tp_sequence_sharded_predicate - if predicate is not None and not predicate(num_tokens_per_dp): + if not sharded: return num_tokens_per_dp, 0 parallel = get_parallel() tokens_per_rank = num_tokens_per_dp // parallel.attn_tp_size @@ -275,13 +261,24 @@ def _attn_tp_local_shard_bounds(num_tokens_per_dp: int) -> Tuple[int, int]: def compute_local_num_token_non_padded( global_num_token_non_padded: torch.Tensor, num_tokens_per_dp: int, + *, + sharded: bool, ) -> torch.Tensor: """Compute local non-padded token count for this attention-TP rank. Converts a global count (across all TP ranks) to a local count for this rank. The "global" scope is within the current DP rank; DP is handled via num_tokens_per_dp. + + ``num_tokens_per_dp`` is the padded bucket width for the DP group, so each rank + owns a contiguous ``chunk = num_tokens_per_dp // attn_tp_size`` slice: the local + count is ``clamp(global - chunk * attn_tp_rank, 0, chunk)``. The padded bucket + (not ``ceil(real / attn_tp_size)``) sets the chunk, so a trailing rank can own + zero real tokens. ``sharded`` False returns the global count unchanged + (replicated). """ - tokens_per_rank, rank_offset = _attn_tp_local_shard_bounds(num_tokens_per_dp) + tokens_per_rank, rank_offset = _attn_tp_local_shard_bounds( + num_tokens_per_dp, sharded=sharded + ) return torch.clamp( global_num_token_non_padded - rank_offset, 0, @@ -292,15 +289,14 @@ def compute_local_num_token_non_padded( def compute_local_num_token_non_padded_cpu( global_num_token_non_padded: int, num_tokens_per_dp: int, + *, + sharded: bool, ) -> int: - """Int-scalar twin of ``compute_local_num_token_non_padded``. - - Replay-time hooks hold the global count as a host int - (``num_token_non_padded_cpu``) and write the localized result into a - device buffer; keeping the math on ints lets them use ``Tensor.fill_`` - instead of staging a CPU tensor through a host-to-device copy per replay. - """ - tokens_per_rank, rank_offset = _attn_tp_local_shard_bounds(num_tokens_per_dp) + """Int-scalar twin of ``compute_local_num_token_non_padded`` for replay-time + hooks that hold the global count as a host int.""" + tokens_per_rank, rank_offset = _attn_tp_local_shard_bounds( + num_tokens_per_dp, sharded=sharded + ) return min(max(global_num_token_non_padded - rank_offset, 0), tokens_per_rank) @@ -529,9 +525,36 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): # Has to be None when cuda graph is captured. global_num_tokens_for_logprob_cpu: Optional[List[int]] = None global_num_tokens_for_logprob_gpu: Optional[torch.Tensor] = None - # For padding - num_token_non_padded: Optional[torch.Tensor] = None # scalar tensor - num_token_non_padded_cpu: int = None + + # Real (non-padding) token count, held at two scopes whose meaning never + # changes once set: + # GLOBAL — the real-token count across the attn-TP group before sharding. + # global_num_token_non_padded GPU int32 scalar. The invariant source + # the eager forward and the cuda-graph + # registry localize from on each forward / + # replay. Present only when + # enable_num_token_non_padded() + # (moe_ep_size > 1). + # global_num_token_non_padded_cpu host int. Host-side attention/backend + # slices read it directly; the prefill + # graph registry derives its per-rank GPU + # scalar from it. + # LOCAL — this attn-TP rank's owned count after sharding. + # num_token_non_padded GPU int32 scalar, derived from + # global_num_token_non_padded (see + # compute_local_num_token_non_padded). The + # MoE topk kernel masks padded rows with + # it; replicated forwards keep the full + # count. Left None until localized (eager + # prep / graph replay). Present only when + # enable_num_token_non_padded(). + global_num_token_non_padded: Optional[torch.Tensor] = None # scalar, GLOBAL + global_num_token_non_padded_cpu: int = None # host int, GLOBAL + num_token_non_padded: Optional[torch.Tensor] = None # scalar, LOCAL (derived) + + # Whether this forward's sequence is sharded across the attn-TP group (SP on) + # vs. replicated; stamped per forward, defaults to replicated. + attn_tp_sequence_sharded: bool = False # === Runtime-filled (set during the forward pass / cuda graph / managers; not at construction) === # Preallocated piecewise-graph attention output, set by RadixAttention. @@ -855,12 +878,12 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): num_tokens = len(batch.input_ids) if batch.input_ids is not None else 0 if enable_num_token_non_padded(): - ret.num_token_non_padded = torch.tensor( + ret.global_num_token_non_padded = torch.tensor( num_tokens, dtype=torch.int32, pin_memory=is_pin_memory_available(device), ).to(device, non_blocking=True) - ret.num_token_non_padded_cpu = num_tokens + ret.global_num_token_non_padded_cpu = num_tokens ret.init_mlp_sync_metadata(batch, device) @@ -1005,21 +1028,29 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): pin_memory=is_pin_memory_available(batch.device), ).to(batch.device, non_blocking=True) - def adjust_num_token_non_padded_for_attn_tp(self) -> None: - """Make num_token_non_padded local to this attention-TP rank.""" + def set_local_num_token_non_padded(self, *, sharded: bool) -> None: + """Derive the LOCAL num_token_non_padded from the invariant GLOBAL scalar. + + A replicated (``sharded=False``) forward keeps the full DP-group count. + """ from sglang.srt.utils.common import require_mlp_tp_gather - dp_rank = get_parallel().attn_dp_rank - assert self.global_num_tokens_cpu is not None - - if require_mlp_tp_gather(): - num_tokens_per_dp = self.global_num_tokens_cpu[dp_rank] + if self.global_num_tokens_cpu is not None: + # DP / MLP-sync path: per-DP padded width. + if require_mlp_tp_gather(): + num_tokens_per_dp = self.global_num_tokens_cpu[ + get_parallel().attn_dp_rank + ] + else: + num_tokens_per_dp = self.global_num_tokens_cpu[0] else: - num_tokens_per_dp = self.global_num_tokens_cpu[0] + # Pure TP+SP: local input width. + num_tokens_per_dp = self._forward_num_tokens() self.num_token_non_padded = compute_local_num_token_non_padded( - global_num_token_non_padded=self.num_token_non_padded, + global_num_token_non_padded=self.global_num_token_non_padded, num_tokens_per_dp=num_tokens_per_dp, + sharded=sharded, ) def merge_mm_inputs(self) -> Optional[MultimodalInputs]: @@ -1369,6 +1400,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): else: num_tokens = global_num_tokens[0] + self.attn_tp_sequence_sharded = model_runner.attn_tp_sequence_sharded( + num_tokens + ) + self.global_dp_buffer_len = buffer_len set_dp_buffer_len( buffer_len, @@ -1456,14 +1491,17 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): mask_dummy_tokens = ( not hybrid_ssm and self._original_forward_mode.is_idle() ) + # Bump the GLOBAL scalar; the LOCAL count is derived from it + # downstream. (global_num_token_non_padded is None unless + # moe_ep_size > 1.) if mask_dummy_tokens: - if self.num_token_non_padded is not None: - self.num_token_non_padded.fill_(0) - self.num_token_non_padded_cpu = 0 + if self.global_num_token_non_padded is not None: + self.global_num_token_non_padded.fill_(0) + self.global_num_token_non_padded_cpu = 0 else: - if self.num_token_non_padded is not None: - self.num_token_non_padded.fill_(num_tokens) - self.num_token_non_padded_cpu = num_tokens + if self.global_num_token_non_padded is not None: + self.global_num_token_non_padded.fill_(num_tokens) + self.global_num_token_non_padded_cpu = num_tokens else: self.extend_num_tokens = bs self.extend_seq_lens = torch.full_like(self.seq_lens, 1) @@ -1633,9 +1671,21 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): spec_info.hidden_states, num_tokens ) + def _forward_num_tokens(self) -> int: + """Token width of this forward (``input_embeds`` over ``input_ids``), as + the model counts it for its SP gate.""" + if self.input_embeds is not None: + return self.input_embeds.shape[0] + return self.input_ids.shape[0] + def prepare_attn_tp_scatter_input(self, model_runner: ModelRunner): from sglang.srt.layers.communicator import get_attn_tp_context + # Pure TP+SP has no MLP-sync pass, so stamp the decision here. + self.attn_tp_sequence_sharded = model_runner.attn_tp_sequence_sharded( + self._forward_num_tokens() + ) + attn_tp_context = get_attn_tp_context() input_scattered = attn_tp_context.use_input_scattered(self) if not input_scattered: diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index cd539390b..353f30878 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -1510,6 +1510,13 @@ class ModelRunner: """Customize a runner-created dummy batch before attention metadata initialization.""" return forward_batch + def attn_tp_sequence_sharded(self, num_tokens: int) -> bool: + """Whether this forward (``num_tokens`` tokens) is attn-TP sharded (SP). + Extension point for per-forward SP gating; the default behavior shards iff + it holds a gathered buffer. + """ + return require_gathered_buffer() + def _prepare_eager_forward_batch(self, forward_batch: ForwardBatch) -> None: """Pad / normalize a batch for the eager (non-cuda-graph) forward. @@ -1524,20 +1531,21 @@ class ModelRunner: else: forward_batch.prepare_attn_tp_scatter_input(self) - # Normalize num_token_non_padded to be local to this attention TP rank if needed. - # The skip is scoped to DSACPLayerCommunicator-style CP (DSA, MLA): those - # flavors already feed a zigzag-split rank-local layout whose token count - # should not be further divided by attn_tp_size. MHA-arch prefill CP - # (Qwen3/Qwen2 MoE) keeps the attn_tp-replicated layout and wants the - # adjustment to run — see docs/design/prefill-cp-mla.md §Phase 5. - if ( - forward_batch.num_token_non_padded is not None - and forward_batch.global_num_tokens_gpu is not None - and require_gathered_buffer() - and not is_dsa_enable_prefill_cp() - and not is_mla_prefill_cp_enabled() - ): - forward_batch.adjust_num_token_non_padded_for_attn_tp() + # Derive the LOCAL num_token_non_padded from the GLOBAL scalar. sharded is + # cleared for DSACPLayerCommunicator-style CP (DSA, MLA): those flavors + # already feed a zigzag-split rank-local layout whose token count should + # not be further divided by attn_tp_size, so they keep the full count. + # MHA-arch prefill CP (Qwen3/Qwen2 MoE) keeps the attn_tp-replicated + # layout and wants sharding to apply — see docs/design/prefill-cp-mla.md + # §Phase 5. + if forward_batch.global_num_token_non_padded is not None: + forward_batch.set_local_num_token_non_padded( + sharded=( + forward_batch.attn_tp_sequence_sharded + and not is_dsa_enable_prefill_cp() + and not is_mla_prefill_cp_enabled() + ), + ) # Hisparse coordinator — backends now read it from self.model_runner. if self.hisparse_coordinator is not None: diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index 82677b901..cb114f8be 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -459,6 +459,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): require_gathered_buffer=self.require_gathered_buffer, enable_prefill_cp=self.enable_prefill_cp, require_mlp_tp_gather=self.require_mlp_tp_gather, + attn_tp_sharded_fn=self.model_runner.attn_tp_sequence_sharded, dp_size=self.dp_size, source=self.buffers, ) @@ -915,17 +916,18 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): else None ) - # Adjust for attention TP if needed (matching replay path in - # populate_from_forward_batch). + # Localize the count when this bucket is attn-TP sharded (SP on). + attn_tp_sharded = self.model_runner.attn_tp_sequence_sharded(num_tokens) buffers.num_token_non_padded[...] = num_tokens if ( enable_num_token_non_padded() - and self.require_gathered_buffer and not self.enable_prefill_cp + and attn_tp_sharded ): local = compute_local_num_token_non_padded( global_num_token_non_padded=buffers.num_token_non_padded, num_tokens_per_dp=num_tokens, + sharded=True, ) buffers.num_token_non_padded.copy_(local) @@ -1009,6 +1011,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): spec_info=spec_info, capture_hidden_mode=self.capture_hidden_mode, num_token_non_padded=buffers.num_token_non_padded, + attn_tp_sequence_sharded=attn_tp_sharded, global_forward_mode=self.capture_forward_mode, lora_ids=lora_ids, rids_int=rids_int, diff --git a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py index 5e2a7f7aa..017a81d69 100644 --- a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py @@ -382,6 +382,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): enable_prefill_cp=( is_dsa_enable_prefill_cp() or is_mla_prefill_cp_enabled() ), + attn_tp_sharded_fn=self.model_runner.attn_tp_sequence_sharded, source=self.buffers, ) @@ -646,10 +647,12 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): buf = self.buffer_registry.get_slot("num_token_non_padded").buffer buf.fill_(num_tokens) - if require_gathered_buffer(): + # Localize the count when this bucket is attn-TP sharded (SP on). + if self.model_runner.attn_tp_sequence_sharded(num_tokens): local = compute_local_num_token_non_padded( global_num_token_non_padded=buf, num_tokens_per_dp=num_tokens, + sharded=True, ) buf.copy_(local) return buf @@ -1390,7 +1393,10 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): # Ported from main #27468. capture_hidden_mode=self.capture_hidden_mode, num_token_non_padded=self._capture_num_token_non_padded(num_tokens), - num_token_non_padded_cpu=num_tokens, + global_num_token_non_padded_cpu=num_tokens, + attn_tp_sequence_sharded=self.model_runner.attn_tp_sequence_sharded( + num_tokens + ), global_forward_mode=ForwardMode.EXTEND, # All-None ids are safe: kernels no-op at rank 0 and replay # refreshes the static batch info with live values. @@ -1660,7 +1666,10 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): spec_info=padded_spec_info, capture_hidden_mode=forward_batch.capture_hidden_mode, num_token_non_padded=num_token_non_padded, - num_token_non_padded_cpu=forward_batch.num_token_non_padded_cpu, + global_num_token_non_padded_cpu=forward_batch.global_num_token_non_padded_cpu, + attn_tp_sequence_sharded=self.model_runner.attn_tp_sequence_sharded( + static_num_tokens + ), global_forward_mode=pcg_global_forward_mode, lora_ids=forward_batch.lora_ids, sampling_info=forward_batch.sampling_info, @@ -1785,10 +1794,20 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): original_layer_forward = self.layer_model.forward self.layer_model.forward = replay_layer_forward - # For Full, run the eager tail against the raw user-facing batch so it - # uses real request metadata instead of padded slots. BCG has no - # request-slot padding, so static_forward_batch is already the serving batch. - tail_batch = forward_batch if full_path else static_forward_batch + if full_path: + # Run the eager tail against the raw user-facing batch so it uses + # real request metadata instead of padded slots. Its SP verdict came + # from the unpadded count, though, while the captured body froze one + # at the padded bucket count; the two can differ, so carry the body's + # on a private view rather than mutating the caller's batch. + tail_batch = copy.copy(forward_batch) + tail_batch.attn_tp_sequence_sharded = ( + static_forward_batch.attn_tp_sequence_sharded + ) + else: + # BCG has no request-slot padding, so static_forward_batch is + # already the serving batch and already holds the body's verdict. + tail_batch = static_forward_batch if not full_path: # MTP consumes the target model's live multimodal embeddings in its # eager wrapper before the captured transformer body is replayed. diff --git a/python/sglang/srt/model_executor/runner_utils/buffers.py b/python/sglang/srt/model_executor/runner_utils/buffers.py index 271c0d250..05d1317ac 100644 --- a/python/sglang/srt/model_executor/runner_utils/buffers.py +++ b/python/sglang/srt/model_executor/runner_utils/buffers.py @@ -329,6 +329,7 @@ class DecodeInputBuffers(ForwardInputBuffers): local = compute_local_num_token_non_padded( global_num_token_non_padded=forward_batch.num_token_non_padded, num_tokens_per_dp=num_tokens_per_dp, + sharded=forward_batch.attn_tp_sequence_sharded, ) dsts.append(self.num_token_non_padded) srcs.append(local) diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 9c279f441..768dcfdc0 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -596,7 +596,7 @@ def deepseek_v4_attention_with_output( forward_batch = context.forward_batch attention_layers = context.attention_layers attention_layer = attention_layers[layer_id] - real_num_tokens = forward_batch.num_token_non_padded_cpu + real_num_tokens = forward_batch.global_num_token_non_padded_cpu if real_num_tokens == 0: output.zero_() diff --git a/python/sglang/srt/models/inkling.py b/python/sglang/srt/models/inkling.py index ae0d280cd..a419de4b4 100644 --- a/python/sglang/srt/models/inkling.py +++ b/python/sglang/srt/models/inkling.py @@ -405,7 +405,7 @@ class InklingDecoderLayer(nn.Module): attn_out / residual_out and returns None (the eager_on_graph copy-back is per-tensor, not per-tuple, so outputs must be pre-allocated buffers).""" forward_batch = get_tc_piecewise_forward_context().forward_batch - n = forward_batch.num_token_non_padded_cpu + n = forward_batch.global_num_token_non_padded_cpu # log_scaling_tau is per-token, so narrow it to match the real tokens too. hs, res = self._attn_block( hidden_states[:n], @@ -429,7 +429,7 @@ class InklingDecoderLayer(nn.Module): """Eager break for the final layer's deferred mlp_sconv: run on the real tokens with the live forward_batch, write the padded output buffer.""" forward_batch = get_tc_piecewise_forward_context().forward_batch - n = forward_batch.num_token_non_padded_cpu + n = forward_batch.global_num_token_non_padded_cpu y = self.mlp_sconv(hidden_states[:n], positions[:n], forward_batch) if self.scattered_sconv: # y is the [n, H/P] shard; the output buffer is post-gather [n, H]. diff --git a/python/sglang/srt/models/nemotron_h_utils.py b/python/sglang/srt/models/nemotron_h_utils.py index b8120f614..3e3240a97 100644 --- a/python/sglang/srt/models/nemotron_h_utils.py +++ b/python/sglang/srt/models/nemotron_h_utils.py @@ -27,9 +27,10 @@ def get_real_num_tokens( ) -> int: """Number of real (non DP-padding) rows in ``hidden_states``.""" real_tokens = hidden_states.shape[0] - num_token_non_padded_cpu = getattr(forward_batch, "num_token_non_padded_cpu", None) - if num_token_non_padded_cpu is not None: - real_tokens = min(real_tokens, int(num_token_non_padded_cpu)) + if forward_batch.global_num_token_non_padded_cpu is not None: + real_tokens = min( + real_tokens, int(forward_batch.global_num_token_non_padded_cpu) + ) if ( forward_batch.forward_mode.is_extend() and not forward_batch.forward_mode.is_mixed() diff --git a/python/sglang/srt/speculative/dspark_components/dspark_draft.py b/python/sglang/srt/speculative/dspark_components/dspark_draft.py index 5aae425b7..39b106aca 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_draft.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_draft.py @@ -426,8 +426,10 @@ class DraftBlockProposer: spec_algorithm=SpeculativeAlgorithm.DSPARK, spec_info=self._draft_block_spec_info, capture_hidden_mode=CaptureHiddenMode.NULL, - num_token_non_padded=_make_num_token_non_padded(draft_num_tokens, device), - num_token_non_padded_cpu=draft_num_tokens, + global_num_token_non_padded=_make_num_token_non_padded( + draft_num_tokens, device + ), + global_num_token_non_padded_cpu=draft_num_tokens, ) self._fill_dp_moe_sync_metadata(draft_forward_batch, batch) graph_runner = self.draft_model_runner.decode_cuda_graph_runner @@ -487,8 +489,8 @@ class DraftBlockProposer: num_tokens = forward_batch.input_ids.numel() num_token_non_padded = _make_num_token_non_padded(num_tokens, device) if num_token_non_padded is not None: - forward_batch.num_token_non_padded = num_token_non_padded - forward_batch.num_token_non_padded_cpu = num_tokens + forward_batch.global_num_token_non_padded = num_token_non_padded + forward_batch.global_num_token_non_padded_cpu = num_tokens forward_batch.global_num_tokens_cpu = gnt forward_batch.global_num_tokens_for_logprob_cpu = gnt_logprob pin_memory = is_pin_memory_available(device) diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_utils.py b/python/sglang/srt/speculative/frozen_kv_mtp_utils.py index c0f48b2b5..af72adf25 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_utils.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_utils.py @@ -129,9 +129,11 @@ def expand_for_topk_draft(forward_batch: ForwardBatch, topk: int) -> None: positions = torch.clamp(forward_batch.seq_lens - 1, min=0).to(torch.int64) forward_batch.positions = positions - forward_batch.num_token_non_padded_cpu = positions.numel() - if forward_batch.num_token_non_padded is not None: - forward_batch.num_token_non_padded.fill_(positions.numel()) + forward_batch.global_num_token_non_padded_cpu = positions.numel() + # Bump the GLOBAL scalar; the LOCAL count is derived from it when the draft + # forward localizes (eager prep / graph replay). + if forward_batch.global_num_token_non_padded is not None: + forward_batch.global_num_token_non_padded.fill_(positions.numel()) if ( forward_batch.mrope_positions is not None and forward_batch.mrope_positions.shape[-1] * topk == positions.numel() diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index 13677db1a..f7b22ef6d 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -354,7 +354,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): extend_seq_lens_cpu=extend_seq_lens_cpu, extend_start_loc=extend_start_loc, extend_num_tokens=self.captured_req_width * bs, - num_token_non_padded_cpu=self.captured_req_width * bs, + global_num_token_non_padded_cpu=self.captured_req_width * bs, return_hidden_states_before_norm=True, ) return forward_batch diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py index 057e8a3db..c42545dda 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py @@ -1629,7 +1629,7 @@ def run_dsa_forward( input_hidden = inputs["input_hidden"] # `input_hidden` may have trailing padding for split-op static-token # contracts; project only the live token rows for QKV. The kernel - # respects `num_token_non_padded_cpu` via the metadata. + # respects `global_num_token_non_padded_cpu` via the metadata. live_input_hidden = input_hidden[: case.num_input_tokens] input_parts = _split_by_lens(live_input_hidden, case.input_lens) kv_hidden = torch.cat( @@ -1668,7 +1668,7 @@ def expected_dsa_output_from_inputs( def dsa_attention_layers(fixture: DSAAttentionFixture) -> list: """Return the RadixAttention layers the backend forwards through. The split-op runner uses this to install per-layer - `num_token_non_padded_cpu` metadata before forward.""" + `global_num_token_non_padded_cpu` metadata before forward.""" return [fixture.actual_module.attn] diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py index 4c2cf225f..8488d04bf 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py @@ -1055,7 +1055,7 @@ def make_lightning_token_padded_inputs( def lightning_attention_layers(fixture: LightningAttentionFixture) -> list: """Return the RadixAttention layers the backend forwards through. The split-op runner uses this list to install per-layer - `num_token_non_padded_cpu` metadata before forward.""" + `global_num_token_non_padded_cpu` metadata before forward.""" return [fixture.actual_module.attn] diff --git a/python/sglang/test/kits/attention_unittest/runner_modes/split_op_runner.py b/python/sglang/test/kits/attention_unittest/runner_modes/split_op_runner.py index c7f44034f..9542b8c96 100644 --- a/python/sglang/test/kits/attention_unittest/runner_modes/split_op_runner.py +++ b/python/sglang/test/kits/attention_unittest/runner_modes/split_op_runner.py @@ -223,13 +223,13 @@ def _make_static_forward_batch(raw_batch, static_num_tokens: int, device: str): dim=0, ) - raw_batch.num_token_non_padded_cpu = raw_num_tokens + raw_batch.global_num_token_non_padded_cpu = raw_num_tokens return replace( raw_batch, input_ids=input_ids, positions=positions, out_cache_loc=out_cache_loc, - num_token_non_padded_cpu=raw_num_tokens, + global_num_token_non_padded_cpu=raw_num_tokens, ) diff --git a/python/sglang/test/kv_canary/fixtures.py b/python/sglang/test/kv_canary/fixtures.py index d8882d418..1125405c0 100644 --- a/python/sglang/test/kv_canary/fixtures.py +++ b/python/sglang/test/kv_canary/fixtures.py @@ -175,7 +175,7 @@ def make_forward_batch( input_ids: Optional[torch.Tensor] = None, positions: Optional[torch.Tensor] = None, out_cache_loc: Optional[torch.Tensor] = None, - num_token_non_padded_cpu: Optional[int] = None, + global_num_token_non_padded_cpu: Optional[int] = None, ) -> SimpleNamespace: seq_lens_default = list(seq_lens_list[:bs]) if req_pool_indices is None: @@ -216,7 +216,7 @@ def make_forward_batch( input_ids=input_ids, positions=positions, out_cache_loc=out_cache_loc, - num_token_non_padded_cpu=num_token_non_padded_cpu, + global_num_token_non_padded_cpu=global_num_token_non_padded_cpu, req_all_ids_flat=None, req_all_ids_lens=None, ) diff --git a/test/registered/attention/unittests/KNOWN_FAILURES.md b/test/registered/attention/unittests/KNOWN_FAILURES.md index eefbf2d5c..04d599f67 100644 --- a/test/registered/attention/unittests/KNOWN_FAILURES.md +++ b/test/registered/attention/unittests/KNOWN_FAILURES.md @@ -174,9 +174,9 @@ moment production is fixed; no test method invokes them today. | Backend | Status | Root cause | |---|---|---| | Lightning | `[no test]` (`lightning/README.md`) | Backend returns flat `[T, num_heads * head_dim]` at `lightning_backend.py:335`; `RadixAttention` piecewise writes per-head (`radix_attention.py:124-137`). Shape mismatch eager vs piecewise | -| Mamba2 | `[no test]` (`mamba/README.md`) | `MambaMixer2.forward` projects ALL rows of `hidden_states` before per-layer `num_token_non_padded_cpu` slicing (`mamba.py:467`); trips assert under token-padding | +| Mamba2 | `[no test]` (`mamba/README.md`) | `MambaMixer2.forward` projects ALL rows of `hidden_states` before per-layer `global_num_token_non_padded_cpu` slicing (`mamba.py:467`); trips assert under token-padding | | DSV4 | `[no test]` (`dsv4/README.md`) | `flash_mla.flash_mla_with_kvcache` asserts `indices.shape == (b, s_q, topk)`; metadata sized for live batch, q is static-token-padded | -| DSA MHA_ONE_SHOT dense fallback | `[no test]` (`dsa/README.md`) | DSA passes K as concatenated `prefix + extend` to `module.attn(save_kv_cache=False)`; `unified_attention_with_output` (`radix_attention.py:170-208`) slices K to `num_token_non_padded_cpu`, dropping the prefix portion — piecewise CG diverges from eager ~50% mismatch (~0.35 max diff) | +| DSA MHA_ONE_SHOT dense fallback | `[no test]` (`dsa/README.md`) | DSA passes K as concatenated `prefix + extend` to `module.attn(save_kv_cache=False)`; `unified_attention_with_output` (`radix_attention.py:170-208`) slices K to `global_num_token_non_padded_cpu`, dropping the prefix portion — piecewise CG diverges from eager ~50% mismatch (~0.35 max diff) | ## C.5. Sparse-kernel production bugs diff --git a/test/registered/attention/unittests/dsa/README.md b/test/registered/attention/unittests/dsa/README.md index cf6a82d90..d629a5578 100644 --- a/test/registered/attention/unittests/dsa/README.md +++ b/test/registered/attention/unittests/dsa/README.md @@ -73,7 +73,7 @@ hardware/SDK. The variant tests live in `test_dsa.py` as `[sum(seq_lens), num_kv_heads, head_dim]`) to `module.attn(q, k, v, forward_batch, save_kv_cache=False)`, but `unified_attention_with_output` (`radix_attention.py:170-208`, which RadixAttention routes to under - piecewise CG) slices K to `forward_batch.num_token_non_padded_cpu` (= + piecewise CG) slices K to `forward_batch.global_num_token_non_padded_cpu` (= live extend-token count) on the per-token K convention used by Triton/FlashInfer/FA. The slice removes the prefix portion, so a piecewise CG run diverges from the eager DSA dense fallback by ~50% diff --git a/test/registered/attention/unittests/dsa/test_dsa.py b/test/registered/attention/unittests/dsa/test_dsa.py index dcd034352..d236cb29c 100644 --- a/test/registered/attention/unittests/dsa/test_dsa.py +++ b/test/registered/attention/unittests/dsa/test_dsa.py @@ -42,7 +42,7 @@ class TestDSAAttentionBackendCorrectness(CustomTestCase): # MHA_ONE_SHOT dense fallback passes K as concatenated prefix+extend # (length = sum(seq_lens)) to `module.attn`, but # `unified_attention_with_output` (`radix_attention.py:170-208`) slices - # K to `forward_batch.num_token_non_padded_cpu` (= live extend-token + # K to `forward_batch.global_num_token_non_padded_cpu` (= live extend-token # count), under the per-token K convention used by Triton/FlashInfer/ # FA. The K-slice removes the prefix portion, so DSA's dense fallback # output diverges by ~50% mismatch under piecewise CG. See diff --git a/test/registered/attention/unittests/mamba/README.md b/test/registered/attention/unittests/mamba/README.md index 1090f7e71..8ea1977cb 100644 --- a/test/registered/attention/unittests/mamba/README.md +++ b/test/registered/attention/unittests/mamba/README.md @@ -83,7 +83,7 @@ call. - **PCG / BCG split-op extend** — `MambaMixer2.forward` asserts `num_actual_tokens == projected_states.shape[0]` (`mamba.py:467`) at the projection step, BEFORE the - `num_token_non_padded_cpu` slicing kicks in at the attention + `global_num_token_non_padded_cpu` slicing kicks in at the attention dispatch. The shared `_run_split_op_extend_case` pads `hidden_states` to a fixed `static_num_tokens` upper bound to exercise the per-layer slicing contract, but Mamba2 trips this diff --git a/test/registered/attention/unittests/mamba/test_mamba2.py b/test/registered/attention/unittests/mamba/test_mamba2.py index 46adfb46d..fa13775c1 100644 --- a/test/registered/attention/unittests/mamba/test_mamba2.py +++ b/test/registered/attention/unittests/mamba/test_mamba2.py @@ -166,9 +166,9 @@ class TestTritonMamba2BackendCorrectness(CustomTestCase): # exactly, with no padding tolerance. The shared split-op runner # pads `hidden_states` to a fixed `static_num_tokens` upper bound # and then relies on the backend's per-layer slicing contract via - # `num_token_non_padded_cpu`. Mamba2 doesn't support this padding + # `global_num_token_non_padded_cpu`. Mamba2 doesn't support this padding # because its mixer projects BEFORE the attention dispatch sees - # `num_token_non_padded_cpu`. Landing this needs either a Mamba2 + # `global_num_token_non_padded_cpu`. Landing this needs either a Mamba2 # mixer change to accept padded `hidden_states`, or a split-op # runner variant that passes unpadded `hidden_states` while still # padding the `forward_batch.input_ids` / `out_cache_loc`. diff --git a/test/registered/cuda_graph/breakable/test_breakable_cuda_graph.py b/test/registered/cuda_graph/breakable/test_breakable_cuda_graph.py index 9ad679fdd..6832403e4 100644 --- a/test/registered/cuda_graph/breakable/test_breakable_cuda_graph.py +++ b/test/registered/cuda_graph/breakable/test_breakable_cuda_graph.py @@ -222,7 +222,7 @@ class TestBreakableCUDAGraphBasic(CustomTestCase): num_tokens = 3 padded_num_tokens = 5 forward_batch = SimpleNamespace( - num_token_non_padded_cpu=num_tokens, + global_num_token_non_padded_cpu=num_tokens, out_cache_loc=torch.arange(padded_num_tokens, device=self.device), positions=torch.arange(padded_num_tokens, device=self.device), ) diff --git a/test/registered/kv_canary/test_self_unit_perturb.py b/test/registered/kv_canary/test_self_unit_perturb.py index 99b4ce19d..10a21dc54 100644 --- a/test/registered/kv_canary/test_self_unit_perturb.py +++ b/test/registered/kv_canary/test_self_unit_perturb.py @@ -206,7 +206,7 @@ class TestRealKvPostForwardPerturb(CustomTestCase): forward_batch.out_cache_loc = torch.tensor( [2], dtype=torch.int32, device=device ) - forward_batch.num_token_non_padded_cpu = 1 + forward_batch.global_num_token_non_padded_cpu = 1 head_snapshot = group.k_head.clone() v_head_snapshot = group.v_head.clone() @@ -288,7 +288,7 @@ class TestReqToTokenPerturb(CustomTestCase): forward_batch.out_cache_loc = torch.tensor( [7, 0, 0], dtype=torch.int32, device=device ) - forward_batch.num_token_non_padded_cpu = 1 + forward_batch.global_num_token_non_padded_cpu = 1 targets = collect_active_slots( maybe_inaccurate_forward_batch=forward_batch, diff --git a/test/registered/kv_canary/test_self_unit_runner_per_forward.py b/test/registered/kv_canary/test_self_unit_runner_per_forward.py index a6e47bbd2..451524e92 100644 --- a/test/registered/kv_canary/test_self_unit_runner_per_forward.py +++ b/test/registered/kv_canary/test_self_unit_runner_per_forward.py @@ -91,7 +91,7 @@ class TestLaunchEndpointsPerForward(CanaryManagerTestCase): forward_batch.out_cache_loc = torch.tensor( [7, 0, 0], dtype=torch.int64, device=self.device ) - forward_batch.num_token_non_padded_cpu = 1 + forward_batch.global_num_token_non_padded_cpu = 1 kernel_launcher_module.launch_endpoints_per_forward( endpoints=(endpoint,), @@ -144,7 +144,7 @@ class TestLaunchEndpointsPerForward(CanaryManagerTestCase): forward_batch.out_cache_loc = torch.tensor( [7], dtype=torch.int32, device=self.device ) - forward_batch.num_token_non_padded_cpu = 1 + forward_batch.global_num_token_non_padded_cpu = 1 kernel_launcher_module.launch_endpoints_per_forward( endpoints=(endpoint,), @@ -182,7 +182,7 @@ class TestLaunchEndpointsPerForward(CanaryManagerTestCase): forward_batch.out_cache_loc = torch.tensor( [7, 0, 0], dtype=torch.int64, device=self.device ) - forward_batch.num_token_non_padded_cpu = 1 + forward_batch.global_num_token_non_padded_cpu = 1 kernel_launcher_module.launch_endpoints_per_forward( endpoints=(endpoint,), @@ -218,7 +218,7 @@ class TestLaunchEndpointsPerForward(CanaryManagerTestCase): forward_batch.out_cache_loc = torch.tensor( [[7, 8]], dtype=torch.int64, device=self.device )[:, 0] - forward_batch.num_token_non_padded_cpu = 1 + forward_batch.global_num_token_non_padded_cpu = 1 kernel_launcher_module.launch_endpoints_per_forward( endpoints=(endpoint,), diff --git a/test/registered/spec/dspark/test_dspark_dp_tier.py b/test/registered/spec/dspark/test_dspark_dp_tier.py index b3278a15b..fea083eef 100644 --- a/test/registered/spec/dspark/test_dspark_dp_tier.py +++ b/test/registered/spec/dspark/test_dspark_dp_tier.py @@ -81,9 +81,11 @@ class TestDraftDpSyncMetadata(CustomTestCase): [1, 3, 0, 2], ) self.assertEqual(forward_batch.global_num_tokens_cpu, [6, 18, 0, 12]) - self.assertEqual(forward_batch.num_token_non_padded.item(), 6) - self.assertEqual(forward_batch.num_token_non_padded.dtype, torch.int32) - self.assertEqual(forward_batch.num_token_non_padded_cpu, 6) + # Metadata fill sets only the invariant GLOBAL count; the LOCAL + # num_token_non_padded is derived later when the draft forward localizes. + self.assertEqual(forward_batch.global_num_token_non_padded.item(), 6) + self.assertEqual(forward_batch.global_num_token_non_padded.dtype, torch.int32) + self.assertEqual(forward_batch.global_num_token_non_padded_cpu, 6) self.assertTrue(forward_batch.can_run_decode_cuda_graph) diff --git a/test/registered/unit/batch_overlap/test_tbo_children_dummy_token_mask.py b/test/registered/unit/batch_overlap/test_tbo_children_dummy_token_mask.py index 2a767b86a..b020ec5da 100644 --- a/test/registered/unit/batch_overlap/test_tbo_children_dummy_token_mask.py +++ b/test/registered/unit/batch_overlap/test_tbo_children_dummy_token_mask.py @@ -6,7 +6,7 @@ by setting ``num_token_non_padded(_cpu)`` to 0 so MoE top-k skips them and attention returns early. ``TboForwardBatchPreparer`` runs right after that, and it used to (a) recompute the children counts from ``len(batch.input_ids)`` -- the *padded* token count, which restores the dummy rows as real tokens for MoE --- and (b) hardcode each child's ``num_token_non_padded_cpu`` to ``None``, so +-- and (b) hardcode each child's ``global_num_token_non_padded_cpu`` to ``None``, so the ``real_num_tokens == 0`` attention skip compared ``None == 0`` and never fired. Net effect: DeepEP/pplx MAX_LEN + ``--enable-two-batch-overlap`` silently lost the whole idle-rank optimization. @@ -29,7 +29,7 @@ from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=5, suite="base-a-test-cpu") -def _make_extend_batch(*, padded_num_tokens: int, num_token_non_padded_cpu: int): +def _make_extend_batch(*, padded_num_tokens: int, global_num_token_non_padded_cpu: int): # Only the fields compute_tbo_children_num_token_non_padded reads. return SimpleNamespace( forward_mode=ForwardMode.EXTEND, @@ -37,7 +37,7 @@ def _make_extend_batch(*, padded_num_tokens: int, num_token_non_padded_cpu: int) tbo_split_seq_index=1, extend_seq_lens_cpu=[4, 4], input_ids=torch.zeros(padded_num_tokens, dtype=torch.long), - num_token_non_padded_cpu=num_token_non_padded_cpu, + global_num_token_non_padded_cpu=global_num_token_non_padded_cpu, ) @@ -49,7 +49,7 @@ def _make_decode_capture_batch(*, num_tokens: int): tbo_split_seq_index=2, extend_seq_lens_cpu=None, input_ids=torch.zeros(num_tokens, dtype=torch.long), - num_token_non_padded_cpu=None, + global_num_token_non_padded_cpu=None, ) @@ -66,30 +66,36 @@ class TestTboChildrenDummyTokenMask(CustomTestCase): def test_masked_idle_parent_yields_zero_token_children(self): # An idle rank masked to 0 real tokens must split into two 0-token # children even though input_ids still holds the padded dummy rows. - batch = _make_extend_batch(padded_num_tokens=16, num_token_non_padded_cpu=0) + batch = _make_extend_batch( + padded_num_tokens=16, global_num_token_non_padded_cpu=0 + ) self.assertEqual(self._children_counts(batch), [0, 0]) def test_padding_is_not_counted_as_real_tokens(self): # A busy rank padded up to a larger bucket must split on its real token # count, not on the padded input_ids length. - batch = _make_extend_batch(padded_num_tokens=16, num_token_non_padded_cpu=8) + batch = _make_extend_batch( + padded_num_tokens=16, global_num_token_non_padded_cpu=8 + ) self.assertEqual(self._children_counts(batch), [4, 4]) def test_cpu_pair_matches_device_pair(self): # prepare() derives the children's CPU counts separately from the device # tensor; the two must not drift apart. - batch = _make_extend_batch(padded_num_tokens=16, num_token_non_padded_cpu=5) + batch = _make_extend_batch( + padded_num_tokens=16, global_num_token_non_padded_cpu=5 + ) cpu_pair = TboForwardBatchPreparer._split_num_token_non_padded( tbo_split_token_index=TboForwardBatchPreparer._compute_split_token_index( batch ), - num_token_non_padded=batch.num_token_non_padded_cpu, + num_token_non_padded=batch.global_num_token_non_padded_cpu, ) self.assertEqual(list(cpu_pair), self._children_counts(batch)) def test_filter_batch_propagates_cpu_count_to_child(self): # Without this the attention 0-token skip (which reads - # num_token_non_padded_cpu) compares None == 0 and never fires. + # global_num_token_non_padded_cpu) compares None == 0 and never fires. bs = 8 parent = ForwardBatch( forward_mode=ForwardMode.TARGET_VERIFY, @@ -118,7 +124,7 @@ class TestTboChildrenDummyTokenMask(CustomTestCase): out_num_token_non_padded=torch.tensor(0), out_num_token_non_padded_cpu=0, ) - self.assertEqual(child.num_token_non_padded_cpu, 0) + self.assertEqual(child.global_num_token_non_padded_cpu, 0) def test_capture_count_falls_back_to_physical_rows(self): batch = _make_decode_capture_batch(num_tokens=8) diff --git a/test/registered/unit/layers/test_radix_attention.py b/test/registered/unit/layers/test_radix_attention.py index a1534e7a2..8cf4cb633 100644 --- a/test/registered/unit/layers/test_radix_attention.py +++ b/test/registered/unit/layers/test_radix_attention.py @@ -69,7 +69,7 @@ class TestRadixAttentionGraphInterface(CustomTestCase): real_num_tokens=2, ): forward_batch = SimpleNamespace( - num_token_non_padded_cpu=real_num_tokens, + global_num_token_non_padded_cpu=real_num_tokens, out_cache_loc=torch.arange(num_tokens, dtype=torch.int64), positions=torch.arange(num_tokens, dtype=torch.int64), _attn_output=None, diff --git a/test/registered/unit/layers/test_radix_linear_attention.py b/test/registered/unit/layers/test_radix_linear_attention.py index e05d185e0..100e1ed7f 100644 --- a/test/registered/unit/layers/test_radix_linear_attention.py +++ b/test/registered/unit/layers/test_radix_linear_attention.py @@ -79,7 +79,7 @@ class TestRadixLinearAttentionPadding(CustomTestCase): original_out_cache_loc = torch.arange(5) forward_batch = SimpleNamespace( forward_mode=_ExtendMode(), - num_token_non_padded_cpu=3, + global_num_token_non_padded_cpu=3, out_cache_loc=original_out_cache_loc, ) @@ -119,7 +119,7 @@ class TestRadixLinearAttentionPadding(CustomTestCase): original_out_cache_loc = torch.arange(5) forward_batch = SimpleNamespace( forward_mode=_TargetVerifyMode(), - num_token_non_padded_cpu=3, + global_num_token_non_padded_cpu=3, out_cache_loc=original_out_cache_loc, ) @@ -158,7 +158,7 @@ class TestRadixLinearAttentionPadding(CustomTestCase): original_out_cache_loc = torch.arange(5) forward_batch = SimpleNamespace( forward_mode=_ExtendMode(), - num_token_non_padded_cpu=3, + global_num_token_non_padded_cpu=3, out_cache_loc=original_out_cache_loc, ) @@ -189,7 +189,7 @@ class TestRadixLinearAttentionPadding(CustomTestCase): with self.subTest(padded_num_tokens=padded_num_tokens): original_out_cache_loc = torch.arange(padded_num_tokens) forward_batch = SimpleNamespace( - num_token_non_padded_cpu=3, + global_num_token_non_padded_cpu=3, out_cache_loc=original_out_cache_loc, ) context = SimpleNamespace( diff --git a/test/registered/unit/model_executor/test_cuda_graph_buffer_registry.py b/test/registered/unit/model_executor/test_cuda_graph_buffer_registry.py index 2bee42d48..f11a35265 100644 --- a/test/registered/unit/model_executor/test_cuda_graph_buffer_registry.py +++ b/test/registered/unit/model_executor/test_cuda_graph_buffer_registry.py @@ -46,7 +46,8 @@ class _MiniForwardBatch: encoder_lens: Optional[torch.Tensor] = None mrope_positions: Optional[torch.Tensor] = None num_token_non_padded: Optional[torch.Tensor] = None - num_token_non_padded_cpu: Optional[int] = None + global_num_token_non_padded: Optional[torch.Tensor] = None + global_num_token_non_padded_cpu: Optional[int] = None global_num_tokens_gpu: Optional[torch.Tensor] = None global_num_tokens_for_logprob_gpu: Optional[torch.Tensor] = None ngram_embedding_info: Optional[object] = None @@ -818,8 +819,9 @@ class TestBuildDecodeRegistry(unittest.TestCase): global_num_tokens_gpu=torch.zeros(1, dtype=torch.int32), global_num_tokens_for_logprob_gpu=torch.zeros(1, dtype=torch.int32), ) - # Gathered (DP) path: post_fill overwrites the FB copy with the local - # count. Pin attn-TP (size=2, rank=0) so the result is deterministic. + # Sharded (SP-on) forward: post_fill derives the LOCAL count from the + # invariant GLOBAL scalar. Pin attn-TP (size=2, rank=0) so the result + # is deterministic. with get_parallel().override(attn_tp_size=2, attn_tp_rank=0): reg = build_decode_registry( device=torch.device("cpu"), @@ -829,18 +831,67 @@ class TestBuildDecodeRegistry(unittest.TestCase): cache_loc_dtype=torch.int64, enable_num_token_non_padded=True, require_gathered_buffer=True, + attn_tp_sharded_fn=lambda num_tokens: True, source=src, ) fb = _MiniForwardBatch( - num_token_non_padded=torch.tensor([100], dtype=torch.int32), + global_num_token_non_padded=torch.tensor([100], dtype=torch.int32), ) reg.fill_from( fb, raw_bs=4, padded_bs=4, raw_num_tokens=4, padded_num_tokens=8 ) # tokens_per_rank = padded_num_tokens(8) // attn_tp_size(2) = 4; - # local = clamp(100 - rank*4, 0, 4) = 4 (NOT the raw FB copy of 100). + # local = clamp(global(100) - rank*4, 0, 4) = 4. self.assertEqual(int(src.num_token_non_padded.item()), 4) + def test_num_token_non_padded_bypass_carries_local_count(self): + # Regression: the dense SBD draft and TBO sub-batches bypass + # ForwardBatch.init_new -- they leave global_num_token_non_padded None and + # set the replicated LOCAL count directly. The decode post_fill must carry + # that value through verbatim, not derive from the absent global (which + # crashed on None - rank_offset). + from sglang.srt.model_executor.cuda_graph_buffer_registry import ( + build_decode_registry, + ) + from sglang.srt.runtime_context import get_parallel + + ntnp = torch.full((1,), 99, dtype=torch.int32) # poisoned static buffer + src = SimpleNamespace( + input_ids=torch.zeros(8, dtype=torch.int64), + positions=torch.zeros(8, dtype=torch.int64), + out_cache_loc=torch.zeros(8, dtype=torch.int64), + req_pool_indices=torch.zeros(4, dtype=torch.int64), + seq_lens=torch.full((4,), 5, dtype=torch.int64), + seq_lens_cpu=torch.full((4,), 5, dtype=torch.int64), + mrope_positions=torch.zeros((3, 8), dtype=torch.int64), + num_token_non_padded=ntnp, + global_num_tokens_gpu=torch.zeros(1, dtype=torch.int32), + global_num_tokens_for_logprob_gpu=torch.zeros(1, dtype=torch.int32), + ) + # attn_tp_sharded_fn=True pins rank 1: were shard math applied it would + # clamp to 3; carrying the local count verbatim proves the bypass + # short-circuits before any sharding. + with get_parallel().override(attn_tp_size=2, attn_tp_rank=1): + reg = build_decode_registry( + device=torch.device("cpu"), + max_bs=4, + max_num_token=8, + seq_len_fill_value=5, + cache_loc_dtype=torch.int64, + enable_num_token_non_padded=True, + require_gathered_buffer=True, + attn_tp_sharded_fn=lambda num_tokens: True, + source=src, + ) + fb = _MiniForwardBatch( + num_token_non_padded=torch.tensor([7], dtype=torch.int32), + global_num_token_non_padded=None, + ) + reg.fill_from( + fb, raw_bs=4, padded_bs=4, raw_num_tokens=4, padded_num_tokens=8 + ) + self.assertEqual(int(src.num_token_non_padded.item()), 7) + def test_register_global_num_tokens_false_carries_fb_values(self): # register_global_num_tokens=False (eager) excludes the computed # global_num_tokens_* slots so the batch's DP values are carried, not @@ -1107,7 +1158,11 @@ class TestBuildPrefillRegistry(unittest.TestCase): self.assertTrue(torch.all(ids[3:8] == 0)) # padded tail reset self.assertTrue(torch.all(ids[8:] == 7)) # beyond the bucket: untouched - def test_num_token_non_padded_scalar_copy(self): + def test_num_token_non_padded_prefill_buffer_adoption(self): + # The prefill num_token_non_padded slot adopts the source's static + # buffer (shared storage), and its post_fill writes the LOCAL count + # derived from the invariant global host int in place — so the static + # buffer exposed by extract_buffer is the same storage. from sglang.srt.model_executor.cuda_graph_buffer_registry import ( build_prefill_registry, ) @@ -1131,9 +1186,11 @@ class TestBuildPrefillRegistry(unittest.TestCase): input_ids=torch.tensor([1, 2, 3], dtype=torch.int64), positions=torch.tensor([4, 5, 6], dtype=torch.int64), out_cache_loc=torch.tensor([8, 9, 10], dtype=torch.int64), - num_token_non_padded=torch.tensor([3], dtype=torch.int32), + global_num_token_non_padded_cpu=3, ) reg.fill_from(fb, raw_bs=1, padded_bs=1, raw_num_tokens=3, padded_num_tokens=8) + # Not sequence-sharded (default predicate): passthrough of the global + # count into the adopted static buffer. self.assertTrue( torch.equal( reg.get_slot("num_token_non_padded").buffer, @@ -1379,11 +1436,16 @@ class TestPrefillNumTokenNonPaddedPostFill(unittest.TestCase): rows of attn-TP rank 0 — REAL tokens — zeroing their MoE output in-graph. The slot's post_fill must instead recompute the local count against ``ctx.padded_num_tokens`` from the batch's un-adjusted global - count (``num_token_non_padded_cpu``), exactly like the decode registry's + count (``global_num_token_non_padded_cpu``), exactly like the decode registry's post_fill does. + + Localization is gated solely on the per-forward sharding decision + (``attn_tp_sharded_fn``): a sharded bucket re-derives the rank-local + count; a replicated one passes the global count through. The cases below + drive that predicate directly via ``sharded``. """ - def _fill(self, *, attn_tp_rank, attn_tp_size, require_gathered_buffer=True): + def _fill(self, *, attn_tp_rank, attn_tp_size, sharded=True, global_count=1018): from unittest import mock from sglang.srt.model_executor.cuda_graph_buffer_registry import ( @@ -1396,14 +1458,14 @@ class TestPrefillNumTokenNonPaddedPostFill(unittest.TestCase): max_num_token=2048, cache_loc_dtype=torch.int64, enable_num_token_non_padded=True, - require_gathered_buffer=require_gathered_buffer, + attn_tp_sharded_fn=lambda num_tokens: sharded, ) # FB tensor carries the RAW-length-localized (stale) value; the CPU # field carries the un-adjusted global count. fb = _MiniForwardBatch( batch_size=1, num_token_non_padded=torch.tensor([509], dtype=torch.int32), - num_token_non_padded_cpu=1018, + global_num_token_non_padded_cpu=global_count, ) with mock.patch( "sglang.srt.model_executor.forward_batch_info.get_parallel", @@ -1431,11 +1493,18 @@ class TestPrefillNumTokenNonPaddedPostFill(unittest.TestCase): # pads. local = clamp(1018 - 512, 0, 512). self.assertEqual(self._fill(attn_tp_rank=1, attn_tp_size=2), 506) - def test_non_gathered_uses_raw_token_count(self): - # Full prefill graphs need the live raw boundary even without a - # gathered buffer so model layers can discard the padded bucket tail. + def test_not_sharded_passes_through_global_count(self): + # A replicated forward owns every row, so the global count is kept. self.assertEqual( - self._fill(attn_tp_rank=0, attn_tp_size=2, require_gathered_buffer=False), + self._fill(attn_tp_rank=0, attn_tp_size=2, sharded=False), + 1018, + ) + + def test_absent_global_count_falls_back_to_raw_tokens(self): + # Full prefill graphs still need the live raw boundary when the batch + # carries no global count, so layers can discard the bucket tail. + self.assertEqual( + self._fill(attn_tp_rank=0, attn_tp_size=2, global_count=None), 1018, ) @@ -1471,10 +1540,10 @@ class TestFillOncePolicy(unittest.TestCase): class TestComputedSlots(unittest.TestCase): - """num_token_non_padded (copy_from_fb + post_fill) and global_num_tokens + """num_token_non_padded and global_num_tokens are both computed slots (copy_from_fb=False + post_fill fill).""" - def test_num_token_non_padded_copy_path(self): + def test_num_token_non_padded_passthrough_path(self): from sglang.srt.model_executor.cuda_graph_buffer_registry import ( build_decode_registry, ) @@ -1491,10 +1560,11 @@ class TestComputedSlots(unittest.TestCase): self.assertTrue(reg.has_slot("num_token_non_padded")) fb = _MiniForwardBatch( batch_size=2, - num_token_non_padded=torch.tensor([7], dtype=torch.int32), + global_num_token_non_padded=torch.tensor([7], dtype=torch.int32), ) - reg.fill_from(fb, raw_bs=2, padded_bs=2, raw_num_tokens=2, padded_num_tokens=2) - # Non-gathered: plain FB copy, post_fill is a no-op. + reg.fill_from(fb, raw_bs=2, padded_bs=2, raw_num_tokens=7, padded_num_tokens=8) + # Not sequence-sharded (default predicate): post_fill passes the global + # scalar (7) through to the local buffer unchanged (clamped to bucket 8). self.assertTrue( torch.equal( reg.get_slot("num_token_non_padded").buffer, diff --git a/test/registered/unit/model_executor/test_num_token_non_padded_localization.py b/test/registered/unit/model_executor/test_num_token_non_padded_localization.py new file mode 100644 index 000000000..24dd96663 --- /dev/null +++ b/test/registered/unit/model_executor/test_num_token_non_padded_localization.py @@ -0,0 +1,70 @@ +"""Per-rank-local non-padded token count across attn x MoE parallelism layouts. + +``compute_local_num_token_non_padded`` (GPU tensor) and +``compute_local_num_token_non_padded_cpu`` (host int) convert a dp-group-global +real-token count into this attention-TP rank's local count. Each rank owns a +contiguous ``padded_bucket // attn_tp_size`` slice of the padded sequence, so the +localizer clamps ``real - chunk * attn_tp_rank`` into ``[0, chunk]``: a replicated +(non-sharded) rank keeps the full count and SP ranks split it. The value is +identical whether the MoE runs TP or EP -- it is an attention-side quantity both +backends consume. This table locks the exact per-rank counts and that the GPU +tensor and host-int twin agree, so a change to the sharding math fails loudly. +""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + +import unittest + +import torch + +from sglang.srt.model_executor.forward_batch_info import ( + compute_local_num_token_non_padded, + compute_local_num_token_non_padded_cpu, +) +from sglang.srt.runtime_context import get_parallel +from sglang.test.test_utils import CustomTestCase + + +class TestNumTokenNonPaddedLayoutTable(CustomTestCase): + # (label, attn_tp_size, sharded, padded_bucket, real-per-dp-group, + # expected [per attn-tp rank] per dp group) + _LAYOUTS = [ + # Each dp rank gets a 10-token request, cuda graph pads it to 16. + ("TP4", 4, False, 16, [10], [[10, 10, 10, 10]]), + ("TP4.SP4", 4, True, 16, [10], [[4, 4, 2, 0]]), + ("DP4", 1, False, 16, [10, 10, 10, 10], [[10], [10], [10], [10]]), + ("TP2.DP2.SP2", 2, True, 16, [10, 10], [[8, 2], [8, 2]]), + # dp0/dp2 get 10 tokens, dp1/dp3 get 20; all padded to 32. + ("DP4.EP4", 1, False, 32, [10, 20, 10, 20], [[10], [20], [10], [20]]), + ("TP2.DP2.SP2.EP2", 2, True, 32, [10, 20], [[10, 0], [16, 4]]), + ] + + def test_layouts_match_expected_per_rank(self): + for label, attn_tp, sharded, bucket, dp_reals, expected in self._LAYOUTS: + for dp_idx, real in enumerate(dp_reals): + for rank in range(attn_tp): + want = expected[dp_idx][rank] + with ( + self.subTest(layout=label, dp=dp_idx, rank=rank), + get_parallel().override( + attn_tp_size=attn_tp, attn_tp_rank=rank + ), + ): + got_cpu = compute_local_num_token_non_padded_cpu( + global_num_token_non_padded=real, + num_tokens_per_dp=bucket, + sharded=sharded, + ) + got_gpu = compute_local_num_token_non_padded( + global_num_token_non_padded=torch.tensor(real), + num_tokens_per_dp=bucket, + sharded=sharded, + ) + self.assertEqual(got_cpu, want) + self.assertEqual(int(got_gpu), want) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py b/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py index d3c3bb1c1..a2935ace9 100644 --- a/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py +++ b/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py @@ -15,6 +15,7 @@ from sglang.srt.model_executor.forward_batch_info import ( ForwardMode, PPProxyTensors, ) +from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner_components.cuda_graph_setup import ( capture_prefill_graph, ) @@ -93,6 +94,19 @@ class _FakeBatchRegistry: class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase): + @patch( + "sglang.srt.model_executor.model_runner.require_gathered_buffer", + return_value=True, + ) + def test_default_attn_tp_sequence_sharded_uses_runtime_predicate( + self, mock_require_gathered_buffer + ): + runner = ModelRunner.__new__(ModelRunner) + runner.server_args = object() + + self.assertTrue(runner.attn_tp_sequence_sharded(num_tokens=4)) + mock_require_gathered_buffer.assert_called_once_with() + def test_low_free_memory_still_captures_prefill_graph(self): eager_runner = object() prefill_runner = object() @@ -198,6 +212,7 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase): runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner) runner.capture_num_tokens = [4] runner.buffer_registry = _FakeBatchRegistry() + runner.model_runner = SimpleNamespace(attn_tp_sequence_sharded=lambda _: False) runner.enable_cp_v2_bcg_capture = False runner._is_full_backend = False runner.backend = SimpleNamespace()