[core] Make spec_v2 seq_lens_cpu optional via backend needs_cpu_seq_lens; Triton opts out (#26128)

Co-authored-by: Qiaolin-Yu <liin1211@outlook.com>
This commit is contained in:
Liangsheng Yin
2026-05-29 13:00:32 -07:00
committed by GitHub
co-authored by Qiaolin-Yu
parent ff8ed7a302
commit 6b5f0d0ccb
14 changed files with 176 additions and 32 deletions
@@ -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."""
@@ -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)
+33 -1
View File
@@ -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())
+8 -2
View File
@@ -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
+14
View File
@@ -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
@@ -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)
@@ -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,
@@ -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:
@@ -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
@@ -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,
+23 -9
View File
@@ -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 = (
@@ -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
@@ -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
+2 -1
View File
@@ -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,