[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:
co-authored by
Qiaolin-Yu
parent
ff8ed7a302
commit
6b5f0d0ccb
@@ -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)
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user