diff --git a/python/sglang/srt/layers/attention/base_attn_backend.py b/python/sglang/srt/layers/attention/base_attn_backend.py index 37a5e62d0..d757ce79a 100644 --- a/python/sglang/srt/layers/attention/base_attn_backend.py +++ b/python/sglang/srt/layers/attention/base_attn_backend.py @@ -18,6 +18,9 @@ if TYPE_CHECKING: class AttentionBackend(ABC): """The base class of attention backends""" + # Opt out only when this backend never reads seq_lens_cpu / seq_lens_sum. + needs_cpu_seq_lens: bool = True + @abstractmethod def init_forward_metadata(self, forward_batch: ForwardBatch): """Init the metadata for a forward pass.""" diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index 5f78d61ca..333df0827 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -76,6 +76,10 @@ class ForwardMetadata: class TritonAttnBackend(AttentionBackend): + # CUDA-graph replay rebuilds metadata from preallocated kv_indptr/kv_indices + # buffers; it never reads seq_lens_cpu / seq_lens_sum. + needs_cpu_seq_lens: bool = False + def __init__( self, model_runner: ModelRunner, @@ -575,11 +579,19 @@ class TritonAttnBackend(AttentionBackend): attn_lse = None elif forward_batch.forward_mode.is_draft_extend(): + # Eager only (CG replay bypasses init); explicit D2H here instead of + # letting torch.empty inside generate_attn_arg_prefill .item() on a + # GPU cumsum tensor. + seq_lens_sum = ( + forward_batch.seq_lens_sum + if forward_batch.seq_lens_sum is not None + else int(forward_batch.seq_lens.sum()) + ) kv_indices, kv_indptr, qo_indptr, custom_mask = ( spec_info.generate_attn_arg_prefill( forward_batch.req_pool_indices, forward_batch.seq_lens, - None, + seq_lens_sum, self.req_to_token, ) ) @@ -1263,6 +1275,8 @@ class TritonMultiStepDraftBackend: draft decoding steps. """ + needs_cpu_seq_lens: bool = False + def __init__( self, model_runner: ModelRunner, @@ -1311,6 +1325,10 @@ class TritonMultiStepDraftBackend: num_seqs = forward_batch.batch_size bs = self.topk * num_seqs seq_lens_sum = forward_batch.seq_lens_sum + if seq_lens_sum is None: + # seq_lens_sum here only slice-clamps a preallocated kv_indices buffer; + # over-estimate is safe. Use a static UB to skip the per-iter .sum().item() D2H. + seq_lens_sum = num_seqs * self.max_context_len generate_draft_decode_kv_indices[ (self.speculative_num_steps, num_seqs, self.topk) diff --git a/python/sglang/srt/managers/overlap_utils.py b/python/sglang/srt/managers/overlap_utils.py index 5f622d4ff..51e4d75f6 100644 --- a/python/sglang/srt/managers/overlap_utils.py +++ b/python/sglang/srt/managers/overlap_utils.py @@ -1,7 +1,7 @@ from __future__ import annotations import os -from typing import TYPE_CHECKING, Optional, Union +from typing import TYPE_CHECKING, Optional, Sequence, Union import torch @@ -9,11 +9,32 @@ from sglang.srt.speculative.spec_utils import spec_need_hidden_states from sglang.srt.utils import is_cuda, is_hip, is_npu if TYPE_CHECKING: + from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.mem_cache.memory_pool import ReqToTokenPool + from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.eagle_info import EagleDraftInput from sglang.srt.speculative.spec_info import SpeculativeAlgorithm + +def decide_needs_cpu_seq_lens( + server_args: "ServerArgs", + attn_backends: Sequence["AttentionBackend"], +) -> bool: + """Whether FutureMap must publish seq_lens_cpu / sum. + + OR over per-backend needs_cpu_seq_lens; force True under TBO / piecewise CG + (they read the CPU mirror outside the backend layer). + """ + if server_args.enable_two_batch_overlap: + # FIXME: support TBO without seq lens cpu value + return True + if not server_args.disable_piecewise_cuda_graph: + # FIXME: support PCG without seq lens cpu value + return True + return any(b.needs_cpu_seq_lens for b in attn_backends) + + _is_cuda = is_cuda() _is_hip = is_hip() _is_npu = is_npu() @@ -88,11 +109,15 @@ class FutureMap: device: torch.device, spec_algo: SpeculativeAlgorithm, req_to_token_pool: ReqToTokenPool, + needs_cpu_seq_lens: bool = True, ): # Bufs indexed by req_pool_idx; slot 0 mirrors KV padding row so # CUDA-graph padded batches (req_pool_idx == 0) are harmless. self.device = device self.spec_algo = spec_algo + # Computed by decide_needs_cpu_seq_lens(); see that helper for the + # full decision (per-backend flag + TBO / piecewise CG overrides). + self.needs_cpu_seq_lens = needs_cpu_seq_lens self.req_pool_size = req_to_token_pool.req_to_token.shape[0] self.output_tokens_buf = ( @@ -213,6 +238,13 @@ class FutureMap: self.publish_ready.wait() batch.seq_lens = self.new_seq_lens_buf[fi] + if not self.needs_cpu_seq_lens: + # GPU gather above is kept (SB.seq_lens must advance each verify); + # skip the .cpu() D2H. Downstream takes the GPU-only path. + batch.seq_lens_cpu = None + batch.seq_lens_sum = None + return + if self.fwd_prepare_d2h_stream is None or self.publish_ready is None: batch.seq_lens_cpu = batch.seq_lens.cpu() # bootstrap / non-CUDA batch.seq_lens_sum = int(batch.seq_lens_cpu.sum()) diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 91a0b7d04..0e8c362be 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -2552,7 +2552,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): self.req_pool_indices = self.req_pool_indices[keep_indices_device] self.req_pool_indices_cpu = self.req_pool_indices_cpu[keep_indices] self.seq_lens = self.seq_lens[keep_indices_device] - self.seq_lens_cpu = self.seq_lens_cpu[keep_indices] self.orig_seq_lens = self.orig_seq_lens[keep_indices_device] self.out_cache_loc = None # Sum is recomputed lazily by ForwardBatch.init_new. @@ -2560,6 +2559,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): if self.input_ids is not None: self.input_ids = self.input_ids[keep_indices_device] + # Optional under no-verify-sync; resolve_seq_lens repopulates before forward. + if self.seq_lens_cpu is not None: + self.seq_lens_cpu = self.seq_lens_cpu[keep_indices] self.mamba_track_indices = None self.mamba_track_mask = None @@ -2606,13 +2608,17 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): [self.req_pool_indices_cpu, other.req_pool_indices_cpu] ) self.seq_lens = torch.cat([self.seq_lens, other.seq_lens]) - self.seq_lens_cpu = torch.cat([self.seq_lens_cpu, other.seq_lens_cpu]) self.orig_seq_lens = torch.cat([self.orig_seq_lens, other.orig_seq_lens]) self.out_cache_loc = None # Sum is recomputed lazily by ForwardBatch.init_new. self.seq_lens_sum = None if self.input_ids is not None: self.input_ids = torch.cat([self.input_ids, other.input_ids]) + # Optional under no-verify-sync; drop the mirror if either side absent. + if self.seq_lens_cpu is None or other.seq_lens_cpu is None: + self.seq_lens_cpu = None + else: + self.seq_lens_cpu = torch.cat([self.seq_lens_cpu, other.seq_lens_cpu]) self.mamba_track_indices = None self.mamba_track_mask = None self.mamba_track_seqlens = None diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 723bd2d00..6a9412062 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -145,6 +145,7 @@ from sglang.srt.managers.io_struct import ( UpdateWeightsFromTensorReqInput, ) from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors +from sglang.srt.managers.overlap_utils import decide_needs_cpu_seq_lens from sglang.srt.managers.prefill_delayer import ( PrefillDelayer, PrefillDelayerSinglePassExecutor, @@ -1146,9 +1147,22 @@ class Scheduler( self.future_map = None return + # Workers not on BaseSpecWorker (e.g. FrozenKVMTPWorker) lack the + # override; fall back to target-only so the helper still produces a + # safe decision (no accidental opt-out for unaudited shapes). + if self.draft_worker is not None: + attn_backends = getattr( + self.draft_worker, + "spec_v2_attn_backends", + (self.tp_worker.model_runner.attn_backend,), + ) + else: + attn_backends = (self.tp_worker.model_runner.attn_backend,) + needs_cpu_seq_lens = decide_needs_cpu_seq_lens(self.server_args, attn_backends) self.future_map = self.spec_algorithm.create_future_map( self.device, self.req_to_token_pool, + needs_cpu_seq_lens=needs_cpu_seq_lens, ) self.batch_record_buf = [None] * 2 self.batch_record_ct = 0 diff --git a/python/sglang/srt/managers/scheduler_components/metrics_reporter.py b/python/sglang/srt/managers/scheduler_components/metrics_reporter.py index 17101633a..f05329d46 100644 --- a/python/sglang/srt/managers/scheduler_components/metrics_reporter.py +++ b/python/sglang/srt/managers/scheduler_components/metrics_reporter.py @@ -44,6 +44,13 @@ LOG_FORWARD_ITERS = envs.SGLANG_LOG_FORWARD_ITERS.get() ENABLE_METRICS_DEVICE_TIMER = envs.SGLANG_ENABLE_METRICS_DEVICE_TIMER.get() +def _decode_total_seq_lens(batch: ScheduleBatch) -> int: + """Sync-free sum of seq_lens for decode metrics.""" + if batch.seq_lens_cpu is not None: + return int(batch.seq_lens_cpu.sum().item()) + return sum(req.seqlen for req in batch.reqs) + + @dataclasses.dataclass class PrefillStats: """Stats for logging prefill batch metrics.""" @@ -430,7 +437,7 @@ class SchedulerMetricsReporter: if tokens == 0: return 0.0, 0.0, 0.0 - total_context = float(batch.seq_lens_cpu.sum().item()) + total_context = float(_decode_total_seq_lens(batch)) flops = ( tokens * self._linear_flops_per_token + self._attn_dot_flops_coeff * total_context @@ -736,7 +743,7 @@ class SchedulerMetricsReporter: self.stats.num_grammar_queue_reqs = len(self.scheduler.grammar_manager) self.stats.gen_throughput = self.last_gen_throughput self.stats.cache_hit_rate = cache_hit_rate - self.stats.decode_sum_seq_lens = batch.seq_lens_cpu.sum().item() + self.stats.decode_sum_seq_lens = _decode_total_seq_lens(batch) # Memory pool usage ratios / Absolute token counts pool_stats.update_scheduler_stats(self.stats) diff --git a/python/sglang/srt/model_executor/cuda_graph_runner.py b/python/sglang/srt/model_executor/cuda_graph_runner.py index 6501b6fe1..2f4ddb5c2 100644 --- a/python/sglang/srt/model_executor/cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/cuda_graph_runner.py @@ -1237,11 +1237,16 @@ class CudaGraphRunner: # FIXME: implicit channel for backends (dsv4) that need forward_batch # in replay metadata prep. Should become a real param on the interface. attn_backend._replay_forward_batch = forward_batch + seq_lens_sum_arg = ( + None + if forward_batch.seq_lens_sum is None + else forward_batch.seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value + ) attn_backend.init_forward_metadata_replay_cuda_graph( bs, buffers.req_pool_indices[:bs], buffers.seq_lens[:bs], - forward_batch.seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value, + seq_lens_sum_arg, buffers.encoder_lens[:bs] if self.is_encoder_decoder else None, self.capture_forward_mode, forward_batch.spec_info, diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 080f722b6..ed18949e1 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -508,8 +508,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): else: seq_lens_cpu = batch.seq_lens_cpu - if batch.seq_lens_sum is None: - batch.seq_lens_sum = int(batch.seq_lens_cpu.sum()) + if batch.seq_lens_sum is None and seq_lens_cpu is not None: + batch.seq_lens_sum = int(seq_lens_cpu.sum()) ret = cls( # Required core inputs @@ -629,14 +629,22 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): if ret.positions is None: ret.positions = clamp_position(batch.seq_lens) else: - assert isinstance(extend_seq_lens, list) - assert isinstance(extend_prefix_lens, list) - ret.extend_seq_lens = torch.tensor(extend_seq_lens, dtype=torch.int32).to( - device, non_blocking=True - ) - ret.extend_prefix_lens = torch.tensor( - extend_prefix_lens, dtype=torch.int32 - ).to(device, non_blocking=True) + if isinstance(extend_seq_lens, list): + # Main path: H2D from host lists; populate *_cpu mirrors. + assert isinstance(extend_prefix_lens, list) + ret.extend_seq_lens = torch.tensor( + extend_seq_lens, dtype=torch.int32 + ).to(device, non_blocking=True) + ret.extend_prefix_lens = torch.tensor( + extend_prefix_lens, dtype=torch.int32 + ).to(device, non_blocking=True) + ret.extend_prefix_lens_cpu = extend_prefix_lens + ret.extend_seq_lens_cpu = extend_seq_lens + else: + # gpu_only: device tensors handed in directly; leave *_cpu unset. + assert isinstance(extend_seq_lens, torch.Tensor) + ret.extend_seq_lens = extend_seq_lens + ret.extend_prefix_lens = extend_prefix_lens ret.extend_num_tokens = batch.extend_num_tokens positions, ret.extend_start_loc = compute_position( model_runner.server_args.attention_backend, @@ -646,8 +654,6 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): ) if ret.positions is None: ret.positions = positions - ret.extend_prefix_lens_cpu = extend_prefix_lens - ret.extend_seq_lens_cpu = extend_seq_lens ret.extend_logprob_start_lens_cpu = extend_logprob_start_lens if model_runner.use_ngram_embedding: diff --git a/python/sglang/srt/speculative/base_spec_worker.py b/python/sglang/srt/speculative/base_spec_worker.py index 9faee6b0e..5db7f6fc6 100644 --- a/python/sglang/srt/speculative/base_spec_worker.py +++ b/python/sglang/srt/speculative/base_spec_worker.py @@ -28,6 +28,12 @@ class BaseSpecWorker(ABC): def draft_worker(self) -> BaseDraftWorker: pass + @property + def spec_v2_attn_backends(self) -> tuple: + """Attn backends touched by spec_v2 forward; OR-ed by decide_needs_cpu_seq_lens. + Default returns target only; subclasses extend with draft backends.""" + return (self.target_worker.model_runner.attn_backend,) + @abstractmethod def clear_cache_pool(self): # TODO: move this abstract method to BaseTpWorker and call through self.model_runner diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index d2992024e..e5a73ffe3 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -534,12 +534,14 @@ class EAGLEDraftExtendCudaGraphRunner: forward_batch.spec_info.num_correct_drafts = buffers.num_correct_drafts[:bs] forward_batch.spec_info.num_accept_tokens = buffers.num_accept_tokens[:bs] + seq_lens_sum = forward_batch.seq_lens_sum + if seq_lens_sum is not None: + seq_lens_sum = seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value self.draft_extend_attn_backend.init_forward_metadata_replay_cuda_graph( bs=bs, req_pool_indices=buffers.req_pool_indices, seq_lens=buffers.seq_lens, - seq_lens_sum=forward_batch.seq_lens_sum - + (bs - raw_bs) * self.seq_len_fill_value, + seq_lens_sum=seq_lens_sum, encoder_lens=None, forward_mode=self.forward_mode, spec_info=forward_batch.spec_info, diff --git a/python/sglang/srt/speculative/eagle_info_v2.py b/python/sglang/srt/speculative/eagle_info_v2.py index a4b5064a3..ad65a9556 100644 --- a/python/sglang/srt/speculative/eagle_info_v2.py +++ b/python/sglang/srt/speculative/eagle_info_v2.py @@ -224,8 +224,10 @@ class EagleDraftInputV2Mixin: draft_model_runner: Any, cuda_graph_runner: Any, ): - seq_lens_cpu_ = batch.seq_lens_cpu - extend_num_tokens = len(batch.seq_lens) * num_draft_tokens + bs = len(batch.seq_lens) + extend_num_tokens = bs * num_draft_tokens + # When seq_lens_cpu is absent, stay on GPU-only path -- no .tolist()/.cpu(). + gpu_only = batch.seq_lens_cpu is None batch.spec_info = self batch.input_ids = predict @@ -235,8 +237,16 @@ class EagleDraftInputV2Mixin: batch.model_config.vocab_size, "v2 prepare_for_extend_to_fill_draft_kvcache input_ids", ) - batch.extend_lens = [num_draft_tokens for _ in range(len(batch.seq_lens))] - batch.prefix_lens = seq_lens_cpu_.tolist() + # init_new requires both list or both Tensor; + # gpu_only emits device tensors to skip H2D. + if gpu_only: + batch.prefix_lens = batch.seq_lens.to(torch.int32) + batch.extend_lens = torch.full( + (bs,), num_draft_tokens, dtype=torch.int32, device=batch.seq_lens.device + ) + else: + batch.prefix_lens = batch.seq_lens_cpu.tolist() + batch.extend_lens = [num_draft_tokens] * bs batch.extend_num_tokens = extend_num_tokens capture_mode = ( CaptureHiddenMode.NULL @@ -253,8 +263,9 @@ class EagleDraftInputV2Mixin: # Forward sees post-write length (draft extend writes num_draft_tokens # slots); mutation stays on forward_batch to preserve SB.seq_lens. forward_batch.seq_lens = forward_batch.seq_lens + num_draft_tokens - forward_batch.seq_lens_cpu = forward_batch.seq_lens_cpu + num_draft_tokens - forward_batch.seq_lens_sum = int(forward_batch.seq_lens_cpu.sum()) + if not gpu_only: + forward_batch.seq_lens_cpu = forward_batch.seq_lens_cpu + num_draft_tokens + forward_batch.seq_lens_sum = int(forward_batch.seq_lens_cpu.sum()) can_cuda_graph = cuda_graph_runner and cuda_graph_runner.can_run(forward_batch) if not batch.forward_mode.is_idle() and not can_cuda_graph: draft_model_runner.attn_backend.init_forward_metadata(forward_batch) @@ -295,10 +306,13 @@ class EagleVerifyInputV2Mixin: batch.mamba_track_mask = None batch.mamba_track_seqlens = None - # Populate seq_lens_cpu/seq_lens_sum on the verify input so that - # TBO's split_spec_info can slice the custom_mask correctly. + # TBO's split_spec_info reads these; no-verify-sync leaves both None. self.seq_lens_cpu = batch.seq_lens_cpu - self.seq_lens_sum = int(batch.seq_lens_cpu.sum()) + self.seq_lens_sum = ( + int(batch.seq_lens_cpu.sum()) + if batch.seq_lens_cpu is not None + else None + ) # Get a forward batch batch.forward_mode = ( diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 8f406ad56..186dd0a1e 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -392,6 +392,19 @@ class EagleDraftWorker(BaseDraftWorker): self.target_worker.model_runner.attn_backend.get_verify_buffers_to_fill_after_draft() ) + # build_tree_kernel uses seq_lens_sum only to size the (non-preallocated) + # tree mask; over-size is safe. Skip per-iter .sum().item() D2H via UB. + seq_lens_sum = batch.seq_lens_sum + if seq_lens_sum is None: + if tree_mask_buf is None: + max_context_len = ( + self.target_worker.model_runner.attn_backend.max_context_len + ) + seq_lens_sum = batch.seq_lens.shape[0] * max_context_len + else: + # tree_mask_buf preallocated -> kernel ignores seq_lens_sum. + seq_lens_sum = 0 + ( tree_mask, position, @@ -405,7 +418,7 @@ class EagleDraftWorker(BaseDraftWorker): top_scores_index, draft_tokens, batch.seq_lens, - batch.seq_lens_sum, + seq_lens_sum, self.topk, self.speculative_num_steps, self.speculative_num_draft_tokens, @@ -780,6 +793,16 @@ class EAGLEWorkerV2(BaseSpecWorker): ) self.adaptive_controller.init_states() + @property + def spec_v2_attn_backends(self) -> tuple: + # Every attn backend a spec_v2 forward touches; consumed by + # decide_needs_cpu_seq_lens to gate the seq_lens_cpu D2H. + return ( + self._target_worker.model_runner.attn_backend, + self._draft_worker.draft_attn_backend, + self._draft_worker.draft_extend_attn_backend, + ) + @property def target_worker(self): return self._target_worker diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index 3e18e81b4..24a57375c 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -665,6 +665,13 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): def draft_worker(self): return self._draft_worker + @property + def spec_v2_attn_backends(self) -> tuple: + return ( + self._target_worker.model_runner.attn_backend, + *self._draft_worker.draft_extend_attn_backend_list, + ) + def clear_cache_pool(self): # allocator and kv cache pool are shared with target worker, which are cleared in scheduler pass diff --git a/python/sglang/srt/speculative/spec_info.py b/python/sglang/srt/speculative/spec_info.py index 65daccd13..02cb0c3fd 100644 --- a/python/sglang/srt/speculative/spec_info.py +++ b/python/sglang/srt/speculative/spec_info.py @@ -122,10 +122,11 @@ class SpeculativeAlgorithm(Enum): self, device: torch.device, req_to_token_pool, + needs_cpu_seq_lens: bool = True, ) -> FutureMap: from sglang.srt.managers.overlap_utils import FutureMap - return FutureMap(device, self, req_to_token_pool) + return FutureMap(device, self, req_to_token_pool, needs_cpu_seq_lens) def build_disagg_draft_input( self,