[Feature] Xiaomi MiMo-V2-Flash day0 support (#15207)

Co-authored-by: 谢学扬 <xiexueyang@xiaomi.com>
Co-authored-by: tz <tangzhen3@xiaomi.com>
Co-authored-by: 李家乐 <lijiale10@xiaomi.com>
Co-authored-by: 张晨 <zhangchen50@xiaomi.com>
Co-authored-by: Shaohui Liu <liushaohui3@xiaomi.com>
Co-authored-by: 王晨 <wangchen77@xiaomi.com>
Co-authored-by: jiangzihan <jiangzihan@xiaomi.com>
Co-authored-by: xiexueyang <xyxie_wangyi@163.com>
Co-authored-by: Linghao Zhang <zhanglinghao@xiaomi.com>
Co-authored-by: ispobock <ispobaoke@gmail.com>
Co-authored-by: Liangsheng Yin <lsyincs@gmail.com>
Co-authored-by: JoyFuture <35593546+JoyFuture@users.noreply.github.com>
Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com>
Co-authored-by: Qiaolin Yu <liin1211@outlook.com>
Co-authored-by: root <root@bj9-ml-g8h20e-k8s-slave106-20251106.alicn.idc.xiaomi.com>
This commit is contained in:
Yingchun Lai
2025-12-19 11:40:07 +08:00
committed by GitHub
co-authored by 谢学扬 tz 李家乐 张晨 Shaohui Liu 王晨 jiangzihan xiexueyang Linghao Zhang ispobock Liangsheng Yin JoyFuture Liangsheng Yin Qiaolin Yu root
parent a0985dd5e5
commit 160a06cab2
38 changed files with 5396 additions and 169 deletions
@@ -127,7 +127,9 @@ class EAGLEDraftExtendCudaGraphRunner:
self.seq_lens = torch.full(
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
)
self.extend_seq_lens = torch.ones((self.max_bs,), dtype=torch.int32)
self.extend_seq_lens = torch.full(
(self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32
)
self.accept_length = torch.full(
(self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32
)
@@ -389,14 +391,16 @@ class EAGLEDraftExtendCudaGraphRunner:
self.seq_lens.fill_(self.seq_len_fill_value)
self.out_cache_loc.zero_()
self.positions.zero_()
self.accept_length.fill_(1)
self.extend_seq_lens.fill_(1)
self.accept_length.fill_(self.num_tokens_per_bs)
self.extend_seq_lens.fill_(self.num_tokens_per_bs)
# Common inputs
self.input_ids[:num_tokens].copy_(forward_batch.input_ids)
self.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
if forward_batch.extend_seq_lens is not None:
self.extend_seq_lens[:raw_bs].copy_(forward_batch.extend_seq_lens)
else:
self.extend_seq_lens[:raw_bs].fill_(self.num_tokens_per_bs)
self.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc)
self.positions[:num_tokens].copy_(forward_batch.positions)
if (
@@ -420,6 +424,16 @@ class EAGLEDraftExtendCudaGraphRunner:
if forward_batch.extend_seq_lens_cpu is not None:
self.extend_seq_lens_cpu[:raw_bs] = forward_batch.extend_seq_lens_cpu
else:
self.extend_seq_lens_cpu[:raw_bs] = [self.num_tokens_per_bs] * raw_bs
if bs > raw_bs:
self.extend_seq_lens_cpu[raw_bs:bs] = [self.num_tokens_per_bs] * (
bs - raw_bs
)
forward_batch.spec_info.extend_seq_lens_cpu = list(
self.extend_seq_lens_cpu[:bs]
)
forward_batch.spec_info.extend_seq_lens_tensor = self.extend_seq_lens[:bs]
if bs != raw_bs:
forward_batch.spec_info.positions = self.positions[:num_tokens]
@@ -10,6 +10,7 @@ import triton.language as tl
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.schedule_batch import ModelWorkerBatch, ScheduleBatch
from sglang.srt.mem_cache.chunk_cache import SWAChunkCache
from sglang.srt.mem_cache.common import (
alloc_paged_token_slots_extend,
alloc_token_slots,
@@ -79,6 +80,12 @@ def assign_draft_cache_locs_page_size_1(
@dataclass
class EagleDraftInputV2Mixin:
def prepare_for_decode(self: EagleDraftInput, batch: ScheduleBatch):
if isinstance(batch.tree_cache, SWAChunkCache):
for req in batch.reqs:
batch.tree_cache.evict_swa(
req, req.seqlen - 1, batch.model_config.attention_chunk_size
)
from sglang.srt.speculative.spec_utils import assign_req_to_token_pool_func
bs = batch.batch_size()
+8 -36
View File
@@ -19,6 +19,7 @@ from sglang.srt.managers.io_struct import UpdateWeightsFromTensorReqInput
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.scheduler import GenerationBatchResult
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.mem_cache.chunk_cache import SWAChunkCache
from sglang.srt.mem_cache.common import (
alloc_paged_token_slots_extend,
alloc_token_slots,
@@ -53,6 +54,7 @@ from sglang.srt.speculative.spec_utils import (
draft_tp_context,
fast_topk,
generate_token_bitmask,
get_last_loc_large_page_size_large_top_k,
load_token_map,
select_top_k_tokens,
)
@@ -366,6 +368,12 @@ class EAGLEWorker(TpModelWorker):
)
def _draft_preprocess_decode(self, batch: ScheduleBatch):
if isinstance(batch.tree_cache, SWAChunkCache):
for req in batch.reqs:
batch.tree_cache.evict_swa(
req, req.seqlen - 1, batch.model_config.attention_chunk_size
)
# Parse args
num_seqs = batch.batch_size()
spec_info = batch.spec_info
@@ -1083,39 +1091,3 @@ def get_last_loc_large_page_size_top_k_1(
prefix_lens,
)
return prefix_lens, seq_lens, last_loc
# Disable torch.compile for this function because it will be
# even slower.
# @torch.compile(dynamic=True)
def get_last_loc_large_page_size_large_top_k(
req_to_token: torch.Tensor,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
speculative_num_steps: int,
topk: int,
page_size: int,
):
prefix_lens = seq_lens
last_page_lens = prefix_lens % page_size
num_new_pages_per_topk = (
last_page_lens + speculative_num_steps + page_size - 1
) // page_size
seq_lens = prefix_lens // page_size * page_size + num_new_pages_per_topk * (
page_size * topk
)
extend_lens = seq_lens - prefix_lens
last_loc = get_last_loc(
req_to_token,
req_pool_indices,
prefix_lens,
)
return (
prefix_lens,
seq_lens,
last_loc,
num_new_pages_per_topk,
extend_lens,
last_page_lens,
)
@@ -0,0 +1,655 @@
# Copyright 2023-2024 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
from __future__ import annotations
import bisect
import logging
import time
from typing import TYPE_CHECKING, Callable
import torch
from sglang.srt.layers.dp_attention import DpPaddingMode, set_dp_buffer_len
from sglang.srt.model_executor.cuda_graph_runner import (
CUDA_GRAPH_CAPTURE_FAILED_MSG,
CudaGraphRunner,
DeepEPCudaGraphRunnerAdapter,
LogitsProcessorOutput,
get_batch_sizes_to_capture,
get_global_graph_memory_pool,
model_capture_mode,
set_global_graph_memory_pool,
set_is_extend_in_batch,
set_torch_compile_config,
)
from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
ForwardBatch,
ForwardMode,
)
from sglang.srt.speculative.eagle_info import EagleDraftInput
from sglang.srt.speculative.mtp_utils import assign_new_state_triton
from sglang.srt.speculative.spec_utils import fast_topk
from sglang.srt.utils import (
get_available_gpu_memory,
require_attn_tp_gather,
require_gathered_buffer,
require_mlp_sync,
require_mlp_tp_gather,
)
if TYPE_CHECKING:
from sglang.srt.speculative.mtp_worker_v2 import MTPDraftWorker
logger = logging.getLogger(__name__)
class MTPDraftExtendCudaGraphRunner:
def __init__(self, mtp_worker: MTPDraftWorker, step: int):
# Parse args
self.step = step
self.mtp_worker = mtp_worker
self.model_runner = model_runner = mtp_worker.mtp_model_runner(self.step)
self.forward_mode = ForwardMode.DRAFT_EXTEND_V2
self.graphs = {}
self.output_buffers = {}
self.enable_torch_compile = model_runner.server_args.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.tp_size = self.model_runner.tp_size
self.dp_size = model_runner.server_args.dp_size
self.enable_pdmux = model_runner.server_args.enable_pdmux
self.speculative_num_steps = model_runner.server_args.speculative_num_steps
self.speculative_num_draft_tokens = (
model_runner.server_args.speculative_num_draft_tokens
)
self.topk = model_runner.server_args.speculative_eagle_topk
self.enable_profile_cuda_graph = (
model_runner.server_args.enable_profile_cuda_graph
)
self.capture_bs, self.compile_bs = get_batch_sizes_to_capture(model_runner)
self.padded_static_len = -1
self.deepep_adapter = DeepEPCudaGraphRunnerAdapter()
# For Attention Backend
self.num_tokens_per_bs = self.speculative_num_steps + 1 + step
self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_bs
self.mtp_worker.draft_extend_attn_backend_list[self.step].init_cuda_graph_state(
self.max_bs, self.max_num_token
)
self.seq_len_fill_value = self.mtp_worker.draft_extend_attn_backend_list[
self.step
].get_cuda_graph_seq_len_fill_value()
def init_buffers_and_capture(
self,
cuda_graph_buffers,
offset,
next_cuda_graph_runner,
):
self.next_cuda_graph_runner = next_cuda_graph_runner
self.seq_lens_cpu = cuda_graph_buffers["seq_lens_cpu"]
self.extend_seq_lens_cpu = [self.num_tokens_per_bs] * self.max_bs
if self.enable_torch_compile:
set_torch_compile_config()
# Graph inputs
with torch.device(self.model_runner.device):
# sliced buffers
# slice according to max_num_token
self.input_ids = cuda_graph_buffers["input_ids"][
offset : offset + self.max_num_token
]
self.out_cache_loc = cuda_graph_buffers["out_cache_loc"][
offset : offset + self.max_num_token
]
self.swa_out_cache_loc = cuda_graph_buffers["swa_out_cache_loc"][
offset : offset + self.max_num_token
]
self.positions = cuda_graph_buffers["positions"][
offset : offset + self.max_num_token
]
# shared states
self.seq_lens = cuda_graph_buffers["seq_lens"]
self.req_pool_indices = cuda_graph_buffers["req_pool_indices"]
self.accept_length = cuda_graph_buffers["accept_length"]
self.extend_seq_lens = torch.full(
(self.max_bs,),
self.num_tokens_per_bs,
dtype=torch.int32,
)
self.extend_start_loc = torch.arange(
0,
self.max_bs * self.num_tokens_per_bs,
step=self.num_tokens_per_bs,
dtype=torch.int32,
)
self.mrope_positions = torch.zeros(
(3, self.max_num_token), dtype=torch.int64
)
self.hidden_states = torch.zeros(
(self.max_num_token, self.model_runner.model_config.hidden_size),
dtype=self.model_runner.dtype,
)
if self.require_gathered_buffer:
if self.require_mlp_tp_gather:
self.global_num_tokens_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32
)
self.global_num_tokens_for_logprob_gpu = torch.zeros(
(self.dp_size,), dtype=torch.int32
)
else:
assert self.require_attn_tp_gather
self.global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
self.global_num_tokens_for_logprob_gpu = torch.zeros(
(1,), dtype=torch.int32
)
else:
self.global_num_tokens_gpu = None
self.global_num_tokens_for_logprob_gpu = None
if hasattr(
self.model_runner.model_config.hf_config, "draft_vocab_size"
): # llama_eagle
vocab_size = self.model_runner.model_config.hf_config.draft_vocab_size
elif hasattr(
self.model_runner.model_config.hf_config, "hot_vocab_size"
): # llama_eagle3
vocab_size = self.model_runner.model_config.hf_config.hot_vocab_size
else:
vocab_size = self.model_runner.model_config.vocab_size
self.next_token_logits_buffer = torch.zeros(
(
(
self.max_bs * self.num_tokens_per_bs
if self.forward_mode == ForwardMode.DRAFT_EXTEND_V2
else self.max_bs
),
vocab_size,
),
dtype=torch.float,
)
# Capture
try:
with model_capture_mode():
self.capture()
except RuntimeError as e:
raise Exception(
f"Capture cuda graph failed: {e}\n{CUDA_GRAPH_CAPTURE_FAILED_MSG}"
)
def can_run(self, forward_batch: ForwardBatch):
if self.require_mlp_tp_gather:
cuda_graph_bs = (
max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_bs
if self.model_runner.spec_algorithm.is_eagle()
else max(forward_batch.global_num_tokens_cpu)
)
else:
cuda_graph_bs = forward_batch.seq_lens.numel()
is_bs_supported = (
cuda_graph_bs in self.graphs
if self.disable_padding
else cuda_graph_bs <= self.max_bs
)
if self.require_mlp_sync:
is_bs_supported = is_bs_supported and forward_batch.can_run_dp_cuda_graph
return is_bs_supported
def _create_graph(self):
return torch.cuda.CUDAGraph()
def _capture_init(self, run_once_fn):
for _ in range(2):
torch.cuda.synchronize()
self.model_runner.tp_group.barrier()
run_once_fn()
def _capture_graph(self, graph, pool, stream, run_once_fn):
with torch.cuda.graph(graph, pool=pool, stream=stream):
out = run_once_fn()
return out
def _replay(self, forward_batch: ForwardBatch):
self.graphs[self.bs].replay()
def capture(self):
CudaGraphRunner.capture(self)
def get_forward_batch(self, bs: int) -> ForwardBatch:
num_tokens = bs * self.num_tokens_per_bs
# Graph inputs
input_ids = self.input_ids[:num_tokens]
req_pool_indices = self.req_pool_indices[:bs]
seq_lens = self.seq_lens[:bs]
seq_lens_cpu = self.seq_lens_cpu[:bs]
extend_seq_lens = self.extend_seq_lens[:bs]
extend_seq_lens_cpu = self.extend_seq_lens_cpu[:bs]
extend_start_loc = self.extend_start_loc[:bs]
accept_length = self.accept_length[:bs]
out_cache_loc = self.out_cache_loc[:num_tokens]
positions = self.positions[:num_tokens]
mrope_positions = self.mrope_positions[:, :num_tokens]
hidden_states = self.hidden_states[:num_tokens]
next_token_logits_buffer = self.next_token_logits_buffer[
: bs if self.forward_mode == ForwardMode.DRAFT_EXTEND else num_tokens
]
if self.require_mlp_tp_gather:
self.global_num_tokens_gpu.copy_(
torch.tensor(
[num_tokens] * self.dp_size,
dtype=torch.int32,
device=self.input_ids.device,
)
)
self.global_num_tokens_for_logprob_gpu.copy_(
torch.tensor(
[num_tokens] * self.dp_size,
dtype=torch.int32,
device=self.input_ids.device,
)
)
global_dp_buffer_len = num_tokens * self.dp_size
elif self.require_attn_tp_gather:
self.global_num_tokens_gpu.copy_(
torch.tensor(
[num_tokens],
dtype=torch.int32,
device=self.input_ids.device,
)
)
self.global_num_tokens_for_logprob_gpu.copy_(
torch.tensor(
[bs],
dtype=torch.int32,
device=self.input_ids.device,
)
)
global_dp_buffer_len = num_tokens
else:
global_dp_buffer_len = None
spec_info = EagleDraftInput(
hidden_states=hidden_states,
accept_length=accept_length,
)
spec_info.positions = None
# Forward batch
forward_batch = ForwardBatch(
forward_mode=self.forward_mode,
batch_size=bs,
input_ids=input_ids,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu,
next_token_logits_buffer=next_token_logits_buffer,
req_to_token_pool=self.model_runner.req_to_token_pool,
token_to_kv_pool=self.model_runner.token_to_kv_pool,
out_cache_loc=out_cache_loc,
seq_lens_sum=seq_lens.sum().item(),
return_logprob=False,
positions=positions,
mrope_positions=mrope_positions,
global_num_tokens_gpu=self.global_num_tokens_gpu,
global_num_tokens_for_logprob_gpu=self.global_num_tokens_for_logprob_gpu,
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
global_dp_buffer_len=global_dp_buffer_len,
spec_algorithm=self.model_runner.spec_algorithm,
spec_info=spec_info,
capture_hidden_mode=CaptureHiddenMode.FULL,
attn_backend=self.mtp_worker.draft_extend_attn_backend_list[self.step],
extend_seq_lens=extend_seq_lens,
extend_seq_lens_cpu=extend_seq_lens_cpu,
padded_static_len=self.padded_static_len,
# added args
extend_start_loc=extend_start_loc,
extend_num_tokens=self.num_tokens_per_bs * bs,
num_token_non_padded_cpu=self.num_tokens_per_bs * bs,
return_hidden_states_before_norm=True,
)
return forward_batch
def capture_one_batch_size(self, bs: int, forward: Callable, stream_idx: int = 0):
graph = self._create_graph()
stream = self.stream
self.deepep_adapter.capture(is_extend_in_batch=True)
num_tokens = bs * self.num_tokens_per_bs
forward_batch = self.get_forward_batch(bs)
self.mtp_worker.draft_extend_attn_backend_list[
self.step
].init_forward_metadata_capture_cuda_graph(
bs=bs,
num_tokens=num_tokens,
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
encoder_lens=None,
forward_mode=self.forward_mode,
spec_info=forward_batch.spec_info,
)
# Run and capture
def run_once():
# Clean intermediate result cache for DP attention
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len(
forward_batch.global_dp_buffer_len,
num_tokens,
forward_batch.dp_padding_mode.is_max_len(),
)
set_is_extend_in_batch(False)
# Backup two fields, which will be modified in-place in `draft_forward`.
output_cache_loc_backup = forward_batch.out_cache_loc
hidden_states_backup = forward_batch.spec_info.hidden_states
ret = self.model_runner.model.forward(
forward_batch.input_ids,
forward_batch.positions,
forward_batch,
)
select_index = (
torch.arange(bs, device=self.model_runner.device)
* (self.speculative_num_draft_tokens + self.step)
+ self.accept_length[:bs]
- 1
+ self.step
)
probs = torch.softmax(ret.next_token_logits[select_index], dim=-1)
ret.topk_p, ret.topk_index = fast_topk(probs, self.topk, dim=-1)
if self.next_cuda_graph_runner is not None:
padding_lens = (
self.speculative_num_draft_tokens - self.accept_length[:bs]
)
assign_new_state_triton(
ret.topk_index,
self.input_ids,
self.positions,
self.hidden_states,
self.out_cache_loc,
self.extend_seq_lens,
self.extend_start_loc,
self.next_cuda_graph_runner.input_ids,
self.next_cuda_graph_runner.positions,
self.next_cuda_graph_runner.hidden_states,
self.next_cuda_graph_runner.out_cache_loc,
self.next_cuda_graph_runner.extend_seq_lens,
self.next_cuda_graph_runner.extend_start_loc,
self.next_cuda_graph_runner.seq_lens,
padding_lens,
forward_batch.batch_size,
self.step,
forward_batch.req_pool_indices,
forward_batch.req_to_token_pool.req_to_token,
self.mtp_worker.req_to_hidden_states_pool,
)
self.next_cuda_graph_runner.swa_out_cache_loc.copy_(
self.model_runner.token_to_kv_pool.translate_loc_from_full_to_swa(
self.next_cuda_graph_runner.out_cache_loc
)
)
forward_batch.out_cache_loc = output_cache_loc_backup
forward_batch.spec_info.hidden_states = hidden_states_backup
return ret
self._capture_init(run_once)
out = self._capture_graph(
graph, get_global_graph_memory_pool(), stream, run_once
)
set_global_graph_memory_pool(graph.pool())
return graph, out
def init_replay_state(
self, forward_batch: ForwardBatch, bs: int, raw_bs: int, num_tokens: int
):
# Common inputs
self.input_ids[:num_tokens].copy_(forward_batch.input_ids)
self.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
if forward_batch.extend_seq_lens is not None:
self.extend_seq_lens[:raw_bs].copy_(forward_batch.extend_seq_lens)
self.extend_start_loc[:raw_bs].copy_(forward_batch.extend_start_loc)
self.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc)
self.positions[:num_tokens].copy_(forward_batch.positions)
if (
forward_batch.spec_info.hidden_states.shape[1]
== self.hidden_states.shape[1]
):
self.hidden_states[:num_tokens].copy_(forward_batch.spec_info.hidden_states)
if forward_batch.spec_info.accept_length is not None:
self.accept_length[:raw_bs].copy_(forward_batch.spec_info.accept_length)
self.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices)
if forward_batch.seq_lens_cpu is not None:
if bs != raw_bs:
self.seq_lens_cpu.fill_(self.seq_len_fill_value)
self.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
if forward_batch.extend_seq_lens_cpu is not None:
self.extend_seq_lens_cpu[:raw_bs] = forward_batch.extend_seq_lens_cpu
def replay(self, forward_batch: ForwardBatch, init_state: bool = True):
assert forward_batch.out_cache_loc is not None
self.deepep_adapter.replay()
# batch_size and num_seqs can be different in case there are finished examples
# in the batch, which will not be counted as num_seqs
raw_bs = forward_batch.batch_size
num_tokens = raw_bs * self.num_tokens_per_bs
# num_tokens = forward_batch.input_ids.shape[0]
if self.require_mlp_tp_gather:
max_batch_size = max(forward_batch.original_global_num_tokens_cpu)
index = bisect.bisect_left(self.capture_bs, max_batch_size)
else:
index = bisect.bisect_left(self.capture_bs, raw_bs)
bs = self.capture_bs[index]
if init_state:
self.init_replay_state(forward_batch, bs, raw_bs, num_tokens)
if self.require_gathered_buffer:
self.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs)
self.global_num_tokens_for_logprob_gpu.fill_(bs * self.num_tokens_per_bs)
forward_batch.spec_info.hidden_states = self.hidden_states[:num_tokens]
forward_batch.spec_info.accept_length = self.accept_length[:bs]
forward_batch.spec_info.num_tokens_per_batch = self.num_tokens_per_bs
forward_batch.spec_info.num_tokens_for_logprob_per_batch = 1
forward_batch.spec_info.positions = self.positions[:num_tokens]
forward_batch.spec_info.extend_seq_lens_tensor = self.extend_seq_lens[:bs]
self.mtp_worker.draft_extend_attn_backend_list[
self.step
].init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=self.req_pool_indices,
seq_lens=self.seq_lens,
seq_lens_sum=forward_batch.seq_lens_sum
+ (bs - raw_bs) * self.seq_len_fill_value,
encoder_lens=None,
forward_mode=self.forward_mode,
spec_info=forward_batch.spec_info,
seq_lens_cpu=self.seq_lens_cpu,
)
# Replay
self.raw_bs = raw_bs
self.bs = bs
self._replay(forward_batch)
out = self.output_buffers[bs]
if self.forward_mode == ForwardMode.DRAFT_EXTEND_V2:
# DRAFT_EXTEND_V2: all tokens calculations whether accepted or not.
unpadding_bs = num_tokens
elif bs != raw_bs:
forward_batch.spec_info.accept_length = self.accept_length[:raw_bs]
unpadding_bs = raw_bs
else:
unpadding_bs = None
if unpadding_bs is not None:
out_copy = out
out = LogitsProcessorOutput(
next_token_logits=out.next_token_logits[:unpadding_bs],
hidden_states=out.hidden_states[:unpadding_bs],
)
out.topk_p = out_copy.topk_p[:raw_bs]
out.topk_index = out_copy.topk_index[:raw_bs]
return out
class MTPMultiStepDraftExtendCudaGraphRunner:
def __init__(self, mtp_worker: MTPDraftWorker):
self.mtp_worker = mtp_worker
self.device = mtp_worker.device
self.gpu_id = mtp_worker.gpu_id
self.speculative_num_steps = mtp_worker.speculative_num_steps
self.draft_extend_attn_backend_list = mtp_worker.draft_extend_attn_backend_list
self.runners = []
self.cuda_graph_buffers = {}
self.seq_len_fill_value = 1
self.max_bs = 1
self.offsets = [0]
self._init_and_capture()
def _init_and_capture(self):
if self.mtp_worker.server_args.disable_cuda_graph:
self.runners = [None] * self.speculative_num_steps
return
self.runners = []
buffer_len_list = []
# 1. Capture loop
for step in range(self.speculative_num_steps):
if self.draft_extend_attn_backend_list[step]:
runner = MTPDraftExtendCudaGraphRunner(self.mtp_worker, step)
self.runners.append(runner)
self.seq_len_fill_value = runner.seq_len_fill_value
self.max_bs = runner.max_bs
buffer_len_list.append(runner.max_num_token)
self.offsets.append(self.offsets[-1] + runner.max_num_token)
else:
self.runners.append(None)
# 2. Allocate buffers
self.cuda_graph_buffers["seq_lens_cpu"] = torch.full(
(self.max_bs,),
self.seq_len_fill_value,
dtype=torch.int32,
)
with torch.device(self.device):
# Sliced buffers
self.cuda_graph_buffers["input_ids"] = torch.zeros(
(self.offsets[-1],), dtype=torch.int64
)
self.cuda_graph_buffers["out_cache_loc"] = torch.ones(
(self.offsets[-1],), dtype=torch.int64
)
self.cuda_graph_buffers["swa_out_cache_loc"] = torch.ones(
(self.offsets[-1],), dtype=torch.int64
)
self.cuda_graph_buffers["positions"] = torch.zeros(
(self.offsets[-1],), dtype=torch.int64
)
# Shared states
self.cuda_graph_buffers["seq_lens"] = torch.full(
(self.max_bs,),
self.seq_len_fill_value,
dtype=torch.int32,
)
self.cuda_graph_buffers["req_pool_indices"] = torch.zeros(
(self.max_bs,), dtype=torch.int32
)
self.cuda_graph_buffers["accept_length"] = torch.full(
(self.max_bs,), 1, dtype=torch.int32
)
for step in range(self.speculative_num_steps - 1, -1, -1):
if self.runners[step] is not None:
tic = time.perf_counter()
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info(
f"Capture draft extend cuda graph begin (step {step}). This can take up to several minutes. avail mem={before_mem:.2f} GB"
)
self.runners[step].init_buffers_and_capture(
self.cuda_graph_buffers,
self.offsets[step],
(
self.runners[step + 1]
if step + 1 < self.speculative_num_steps
else None
),
)
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info(
f"Capture draft extend cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB."
)
def reset_buffers(self, forward_batch, batch_result):
self.cuda_graph_buffers["input_ids"].zero_()
self.cuda_graph_buffers["seq_lens"].fill_(self.seq_len_fill_value)
self.cuda_graph_buffers["out_cache_loc"].zero_()
self.cuda_graph_buffers["swa_out_cache_loc"].zero_()
self.cuda_graph_buffers["positions"].zero_()
self.cuda_graph_buffers["accept_length"][: forward_batch.batch_size].copy_(
batch_result.accept_lens
)
def get_runner(self, step):
return self.runners[step]
def get_last_runner(self):
return self.runners[-1] if self.runners else None
def can_run(self, forward_batch):
return self.runners[0].can_run(forward_batch)
+350
View File
@@ -0,0 +1,350 @@
# Copyright 2023-2024 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
import torch
import triton
import triton.language as tl
@triton.jit
def rotate_input_ids_kernel(
input_ids_ptr,
extend_start_loc_ptr,
extend_seq_lens_ptr,
topk_index_ptr,
select_index_ptr,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0)
start_loc = tl.load(extend_start_loc_ptr + pid)
seq_len = tl.load(extend_seq_lens_ptr + pid)
new_token = tl.load(topk_index_ptr + pid)
num_elements_to_shift = seq_len - 1
for off in range(0, num_elements_to_shift, BLOCK_SIZE):
offsets = off + tl.arange(0, BLOCK_SIZE)
mask = offsets < num_elements_to_shift
read_ptr = input_ids_ptr + start_loc + offsets + 1
val = tl.load(read_ptr, mask=mask)
tl.debug_barrier()
write_ptr = input_ids_ptr + start_loc + offsets
tl.store(write_ptr, val, mask=mask)
tl.debug_barrier()
if seq_len > 0:
if select_index_ptr is not None:
last_pos_ptr = input_ids_ptr + tl.load(select_index_ptr + pid)
else:
last_pos_ptr = input_ids_ptr + start_loc + seq_len - 1
tl.store(last_pos_ptr, new_token)
def rotate_input_ids_triton(
input_ids, extend_start_loc, extend_seq_lens, topk_index, select_index=None
):
batch_size = extend_seq_lens.shape[0]
BLOCK_SIZE = 4096 if select_index is not None else 8
grid = (batch_size,)
rotate_input_ids_kernel[grid](
input_ids,
extend_start_loc,
extend_seq_lens,
topk_index,
select_index,
BLOCK_SIZE=BLOCK_SIZE,
)
return input_ids
@triton.jit
def assign_new_state_kernel(
# Source pointers
old_input_ids_ptr,
old_positions_ptr,
old_hidden_states_ptr,
old_out_cache_loc_ptr,
old_extend_seq_lens_ptr,
old_extend_start_loc_ptr,
# Destination pointers
input_ids_ptr,
positions_ptr,
hidden_states_ptr,
out_cache_loc_ptr,
extend_seq_lens_ptr,
extend_start_loc_ptr,
# Auxiliary data pointers
next_token_ids_ptr,
seq_lens_ptr,
padding_lens_ptr,
req_pool_indices_ptr,
req_to_token_ptr,
req_to_hidden_states_pool_ptr,
# Scalars and Strides
step,
stride_hidden_seq,
stride_hidden_dim, # hidden_states strides
stride_pool_req,
stride_pool_step,
stride_pool_dim, # pool strides
stride_req_token_0,
stride_req_token_1, # req_to_token strides
# Meta-parameters
HIDDEN_DIM: tl.constexpr,
BLOCK_SEQ: tl.constexpr,
BLOCK_HID: tl.constexpr,
):
pid = tl.program_id(0)
seq_len: tl.tensor = tl.load(seq_lens_ptr + pid)
old_extend_len = tl.load(old_extend_seq_lens_ptr + pid)
old_start = tl.load(old_extend_start_loc_ptr + pid)
new_extend_len = old_extend_len + 1
new_start = old_start + pid
tl.store(extend_seq_lens_ptr + pid, new_extend_len)
tl.store(extend_start_loc_ptr + pid, new_start)
offs_seq = tl.arange(0, BLOCK_SEQ)
mask_seq = offs_seq < old_extend_len
old_ids = tl.load(old_input_ids_ptr + old_start + offs_seq, mask=mask_seq)
tl.store(input_ids_ptr + new_start + offs_seq, old_ids, mask=mask_seq)
padding_len = tl.load(padding_lens_ptr + pid)
tl.store(
input_ids_ptr + new_start + old_extend_len - padding_len,
tl.load(next_token_ids_ptr + pid),
)
old_pos = tl.load(old_positions_ptr + old_start + offs_seq, mask=mask_seq)
tl.store(positions_ptr + new_start + 1 + offs_seq, old_pos, mask=mask_seq)
tl.store(
positions_ptr + new_start, max(tl.load(old_positions_ptr + old_start) - 1, 0)
)
old_cache = tl.load(old_out_cache_loc_ptr + old_start + offs_seq, mask=mask_seq)
tl.store(out_cache_loc_ptr + new_start + 1 + offs_seq, old_cache, mask=mask_seq)
req_idx = tl.load(req_pool_indices_ptr + pid)
token_idx_col = seq_len - old_extend_len - 1
if token_idx_col >= 0:
req_token_ptr_loc = (
req_to_token_ptr
+ (req_idx * stride_req_token_0)
+ (token_idx_col * stride_req_token_1)
)
last_cache_loc = tl.load(req_token_ptr_loc)
tl.store(out_cache_loc_ptr + new_start, last_cache_loc)
pool_vec_offset_base = ((req_idx + 1) * stride_pool_req) + (
-(step + 1) * stride_pool_step
)
for off_h in range(0, HIDDEN_DIM, BLOCK_HID):
offs_h = off_h + tl.arange(0, BLOCK_HID)
mask_h = offs_h < HIDDEN_DIM
for i in range(BLOCK_SEQ):
if i < old_extend_len:
old_h_ptr = (
old_hidden_states_ptr
+ (old_start + i) * stride_hidden_seq
+ (offs_h * stride_hidden_dim)
)
new_h_ptr = (
hidden_states_ptr
+ (new_start + 1 + i) * stride_hidden_seq
+ (offs_h * stride_hidden_dim)
)
chunk_old = tl.load(old_h_ptr, mask=mask_h)
tl.store(new_h_ptr, chunk_old, mask=mask_h)
pool_ptrs = (
req_to_hidden_states_pool_ptr
+ pool_vec_offset_base
+ (offs_h * stride_pool_dim)
)
pool_val = tl.load(pool_ptrs, mask=mask_h)
new_h_start_ptrs = (
hidden_states_ptr
+ (new_start * stride_hidden_seq)
+ (offs_h * stride_hidden_dim)
)
tl.store(new_h_start_ptrs, pool_val, mask=mask_h)
def assign_new_state_triton(
next_token_ids: torch.Tensor,
old_input_ids: torch.Tensor,
old_positions: torch.Tensor,
old_hidden_states: torch.Tensor,
old_out_cache_loc: torch.Tensor,
old_extend_seq_lens: torch.Tensor,
old_extend_start_loc: torch.Tensor,
input_ids: torch.Tensor,
positions: torch.Tensor,
hidden_states: torch.Tensor,
out_cache_loc: torch.Tensor,
extend_seq_lens: torch.Tensor,
extend_start_loc: torch.Tensor,
seq_lens: torch.Tensor,
padding_lens: torch.Tensor,
num_seqs: int,
step: int,
req_pool_indices: torch.Tensor,
req_to_token: torch.Tensor,
req_to_hidden_states_pool: torch.Tensor,
):
"""
Wrapper function to calculate offsets and launch the Triton kernel.
"""
hidden_dim = hidden_states.shape[1]
BLOCK_SEQ = 8
BLOCK_HID = 64
grid = (num_seqs,)
assign_new_state_kernel[grid](
# Pointers
old_input_ids,
old_positions,
old_hidden_states,
old_out_cache_loc,
old_extend_seq_lens,
old_extend_start_loc,
input_ids,
positions,
hidden_states,
out_cache_loc,
extend_seq_lens,
extend_start_loc,
next_token_ids,
seq_lens,
padding_lens,
req_pool_indices,
req_to_token,
req_to_hidden_states_pool,
# Constants/Strides
step,
old_hidden_states.stride(0),
old_hidden_states.stride(1),
req_to_hidden_states_pool.stride(0),
req_to_hidden_states_pool.stride(1),
req_to_hidden_states_pool.stride(2),
req_to_token.stride(0),
req_to_token.stride(1),
# Meta
HIDDEN_DIM=hidden_dim,
BLOCK_SEQ=BLOCK_SEQ,
BLOCK_HID=BLOCK_HID,
)
@triton.jit
def assign_hidden_states_pool_kernel(
hidden_states_ptr,
req_pool_indices_ptr,
req_to_hidden_states_pool_ptr,
extend_seq_lens_ptr,
extend_start_loc_ptr,
stride_hidden_seq,
stride_hidden_dim,
stride_pool_req,
stride_pool_step,
stride_pool_dim,
HIDDEN_DIM: tl.constexpr,
pool_size: tl.constexpr,
BLOCK_HID: tl.constexpr,
):
pid = tl.program_id(0)
extend_len = tl.load(extend_seq_lens_ptr + pid)
start_loc = tl.load(extend_start_loc_ptr + pid)
end_loc = start_loc + extend_len
req_idx = tl.load(req_pool_indices_ptr + pid)
pool_vec_offset_base = req_idx * stride_pool_req
for i in range(pool_size):
for off_h in range(0, HIDDEN_DIM, BLOCK_HID):
offs_h = off_h + tl.arange(0, BLOCK_HID)
mask_h = offs_h < HIDDEN_DIM
hid_ptr = (
hidden_states_ptr
+ (end_loc - pool_size + i) * stride_hidden_seq
+ offs_h * stride_hidden_dim
)
hid_val = tl.load(hid_ptr, mask=mask_h)
pool_ptr = (
req_to_hidden_states_pool_ptr
+ pool_vec_offset_base
+ i * stride_pool_step
+ offs_h * stride_pool_dim
)
tl.store(pool_ptr, hid_val, mask=mask_h)
def assign_hidden_states_pool_triton(
hidden_states: torch.Tensor,
req_pool_indices: torch.Tensor,
req_to_hidden_states_pool: torch.Tensor,
pool_size: int,
num_seqs: int,
extend_seq_lens: torch.Tensor,
extend_start_loc: torch.Tensor,
):
grid = (num_seqs,)
assign_hidden_states_pool_kernel[grid](
hidden_states,
req_pool_indices,
req_to_hidden_states_pool,
extend_seq_lens,
extend_start_loc,
hidden_states.stride(0),
hidden_states.stride(1),
req_to_hidden_states_pool.stride(0),
req_to_hidden_states_pool.stride(1),
req_to_hidden_states_pool.stride(2),
HIDDEN_DIM=hidden_states.shape[1],
pool_size=pool_size,
BLOCK_HID=64,
)
def assign_hidden_states_pool_torch(
hidden_states: torch.Tensor,
req_pool_indices: torch.Tensor,
req_to_hidden_states_pool: torch.Tensor,
pool_size: int,
num_seqs: int,
extend_seq_lens: torch.Tensor,
extend_start_loc: torch.Tensor,
):
for req in range(num_seqs):
pool_idx = req_pool_indices[req]
extend_len = extend_seq_lens[req]
start_loc = extend_start_loc[req]
end_loc = start_loc + extend_len
req_to_hidden_states_pool[pool_idx, :pool_size, :].copy_(
hidden_states[end_loc - pool_size : end_loc, :]
)
+989
View File
@@ -0,0 +1,989 @@
# Copyright 2023-2024 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
import logging
import time
from typing import List, Optional, Tuple
import torch
from sglang.srt.distributed import get_tp_group
from sglang.srt.layers.dp_attention import get_attention_tp_group
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.moe.utils import speculative_moe_backend_context
from sglang.srt.layers.sampler import get_token_ids_logprobs, get_top_logprobs
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.scheduler import GenerationBatchResult
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.mem_cache.chunk_cache import SWAChunkCache
from sglang.srt.mem_cache.common import (
alloc_paged_token_slots_extend,
alloc_token_slots,
)
from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
ForwardBatch,
ForwardMode,
)
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.draft_utils import DraftBackendFactory
from sglang.srt.speculative.eagle_info import (
EagleDraftInput,
EagleVerifyInput,
EagleVerifyOutput,
)
from sglang.srt.speculative.eagle_utils import (
build_tree_kernel_efficient,
organize_draft_results,
)
from sglang.srt.speculative.eagle_worker import get_last_loc_large_page_size_top_k_1
from sglang.srt.speculative.mtp_draft_extend_cuda_graph_runner import (
MTPDraftExtendCudaGraphRunner,
)
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import (
assign_draft_cache_locs,
detect_nan,
draft_tp_context,
fast_topk,
generate_token_bitmask,
get_last_loc_large_page_size_large_top_k,
load_token_map,
select_top_k_tokens,
)
from sglang.srt.utils import (
empty_context,
get_available_gpu_memory,
get_bool_env_var,
is_cuda,
is_npu,
next_power_of_2,
)
_is_npu = is_npu()
if is_cuda():
from sgl_kernel import segment_packbits # noqa: F401
logger = logging.getLogger(__name__)
SGLANG_RETURN_ORIGINAL_LOGPROB = get_bool_env_var("SGLANG_RETURN_ORIGINAL_LOGPROB")
class MTPWorker(TpModelWorker):
def __init__(
self,
server_args: ServerArgs,
gpu_id: int,
tp_rank: int,
dp_rank: Optional[int],
moe_ep_rank: int,
nccl_port: int,
target_worker: TpModelWorker,
):
# Parse arguments
self.server_args = server_args
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.enable_nan_detection = server_args.enable_nan_detection
self.gpu_id = gpu_id
self.device = server_args.device
self.target_worker = target_worker
self.page_size = server_args.page_size
self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm
)
self.draft_extend_attn_backend_list = []
# Override the context length of the draft model to be the same as the target model.
server_args.context_length = target_worker.model_runner.model_config.context_len
# Do not capture cuda graph in `super().__init__()`
# It will be captured later.
backup_disable_cuda_graph = server_args.disable_cuda_graph
server_args.disable_cuda_graph = True
# Share the allocator with a target worker.
# Draft and target worker own their own KV cache pools.
self.req_to_token_pool, self.token_to_kv_pool_allocator = (
target_worker.get_memory_pool()
)
# Load hot token ids
if self.speculative_algorithm.is_eagle3():
if server_args.speculative_token_map is not None:
logger.warning(
"Speculative token map specified, but EAGLE3 models already have this. Ignoring the specified token map."
)
self.hot_token_id = None
elif server_args.speculative_token_map is not None:
self.hot_token_id = load_token_map(server_args.speculative_token_map)
server_args.json_model_override_args = (
f'{{"hot_vocab_size": {len(self.hot_token_id)}}}'
)
else:
self.hot_token_id = None
# Init draft worker
if server_args.enable_dp_attention and self.speculative_algorithm.is_eagle3():
ctx = draft_tp_context(get_attention_tp_group())
else:
ctx = empty_context()
with ctx, speculative_moe_backend_context():
super().__init__(
server_args=server_args,
gpu_id=gpu_id,
tp_rank=tp_rank,
pp_rank=0, # FIXME
dp_rank=dp_rank,
moe_ep_rank=moe_ep_rank,
nccl_port=nccl_port,
is_draft_worker=True,
req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
is_mtp_worker=True,
)
embed, head = self.target_worker.model_runner.model.get_embed_and_head()
if self.speculative_algorithm.is_eagle3():
# most cases EAGLE3 models don't share lm_head
# but some models (e.g. nvidia/gpt-oss-120b-Eagle3) shares
if (
hasattr(self.draft_model_runner.model, "load_lm_head_from_target")
and self.draft_model_runner.model.load_lm_head_from_target
):
self.draft_model_runner.model.set_embed_and_head(embed, head)
else:
self.draft_model_runner.model.set_embed(embed)
# grab hot token ids
if self.draft_model_runner.model.hot_token_id is not None:
self.hot_token_id = self.draft_model_runner.model.hot_token_id.to(
embed.device
)
else:
if self.hot_token_id is not None:
head = head.clone()
self.hot_token_id = self.hot_token_id.to(head.device)
head.data = head.data[self.hot_token_id]
# Share the embedding and lm_head
for i in range(self.speculative_num_steps):
self.mtp_model_runner(i).model.set_embed_and_head(embed, head)
# Init attention backend and cuda graphs
for i in range(self.speculative_num_steps):
self.mtp_model_runner(i).server_args.disable_cuda_graph = (
backup_disable_cuda_graph
)
self.draft_tp_context = (
draft_tp_context if server_args.enable_dp_attention else empty_context
)
with self.draft_tp_context(
self.mtp_model_runner(0).tp_group
), speculative_moe_backend_context():
self.init_attention_backend()
self.init_cuda_graphs()
# Some dummy tensors
self.num_new_pages_per_topk = torch.empty(
(), dtype=torch.int64, device=self.device
)
self.extend_lens = torch.empty((), dtype=torch.int64, device=self.device)
def init_attention_backend(self):
# Create multi-step attn backends and cuda graph runners
for step in range(self.speculative_num_steps):
draft_backend_factory = DraftBackendFactory(
self.server_args,
self.mtp_model_runner(step),
self.topk,
self.speculative_num_steps,
)
# Initialize draft extend attention backend (respects speculative_attention_mode setting)
self.draft_extend_attn_backend_list.append(
draft_backend_factory.create_draft_extend_backend()
)
def init_cuda_graphs(self):
"""Capture cuda graphs."""
self.cuda_graph_runner_for_draft_extend_list = []
if self.server_args.disable_cuda_graph:
return
# Capture extend
for step in range(self.speculative_num_steps):
if self.draft_extend_attn_backend_list[step] and not _is_npu:
tic = time.perf_counter()
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info(
f"Capture draft extend cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB"
)
self.cuda_graph_runner_for_draft_extend_list.append(
MTPDraftExtendCudaGraphRunner(self, step)
)
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info(
f"Capture draft extend cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB."
)
def mtp_model_runner(self, layer_id: int):
return self.model_runner_list[layer_id]
def forward_batch_generation(self, batch: ScheduleBatch) -> GenerationBatchResult:
"""Run speculative decoding forward.
NOTE: Many states of batch is modified as you go through. It is not guaranteed that
the final output batch have the same state as the input.
Args:
batch: The batch to run forward. The state of the batch is modified as it runs.
Returns:
A tuple of the final logit output of the target model, next tokens accepted,
the batch id (used for overlap schedule), and number of accepted tokens.
"""
if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
logits_output, next_token_ids, seq_lens_cpu = self.forward_target_extend(
batch
)
with self.draft_tp_context(
self.mtp_model_runner(0).tp_group
), speculative_moe_backend_context():
self.forward_draft_extend(
batch, logits_output.hidden_states, next_token_ids, seq_lens_cpu
)
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=next_token_ids,
num_accepted_tokens=0,
can_run_cuda_graph=False,
)
else:
with self.draft_tp_context(
self.mtp_model_runner(0).tp_group
), speculative_moe_backend_context():
spec_info = self.draft(batch)
logits_output, verify_output, model_worker_batch, can_run_cuda_graph = (
self.verify(batch, spec_info)
)
with self.draft_tp_context(
self.mtp_model_runner(0).tp_group
), speculative_moe_backend_context():
# NOTE: We should use `check_forward_draft_extend_after_decode`
# when DP attention is enabled, but it is slow. Skip it for now.
if (
self.server_args.enable_dp_attention
or batch.spec_info.verified_id.shape[0] > 0
):
# decode is not finished
self.forward_draft_extend_after_decode(batch)
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=verify_output.verified_id,
num_accepted_tokens=sum(verify_output.accept_length_per_req_cpu),
can_run_cuda_graph=can_run_cuda_graph,
)
def check_forward_draft_extend_after_decode(self, batch: ScheduleBatch):
local_need_forward = batch.spec_info.verified_id.shape[0] > 0
if not self.server_args.enable_dp_attention:
return local_need_forward
global_need_forward = torch.tensor(
[
(local_need_forward),
],
dtype=torch.int64,
)
torch.distributed.all_reduce(
global_need_forward, group=get_tp_group().cpu_group
)
global_need_forward_cnt = global_need_forward[0].item()
need_forward = global_need_forward_cnt > 0
return need_forward
def forward_target_extend(
self, batch: ScheduleBatch
) -> Tuple[LogitsProcessorOutput, torch.Tensor, int, Optional[torch.Tensor]]:
"""Run the target extend.
Args:
batch: The batch to run. States could be modified.
Returns:
logits_output: The output of logits. It will contain the full hidden states.
next_token_ids: Next token ids generated.
"""
# Forward with the target model and get hidden states.
# We need the full hidden states to prefill the KV cache of the draft model.
model_worker_batch = batch.get_model_worker_batch()
model_worker_batch.capture_hidden_mode = CaptureHiddenMode.FULL
model_worker_batch.return_hidden_states_before_norm = True
batch_result = self.target_worker.forward_batch_generation(model_worker_batch)
logits_output, next_token_ids = (
batch_result.logits_output,
batch_result.next_token_ids,
)
return (
logits_output,
next_token_ids,
model_worker_batch.seq_lens_cpu,
)
def _draft_preprocess_decode(self, batch: ScheduleBatch):
if isinstance(batch.tree_cache, SWAChunkCache):
for req in batch.reqs:
batch.tree_cache.evict_swa(
req, req.seqlen - 1, batch.model_config.attention_chunk_size
)
# Parse args
num_seqs = batch.batch_size()
spec_info = batch.spec_info
# Accumulate penalty
if batch.sampling_info.penalizer_orchestrator.is_required:
# This is a relaxed version of penalties for speculative decoding.
batch.sampling_info.penalizer_orchestrator.cumulate_output_tokens(
spec_info.verified_id.to(torch.int64)
)
# Allocate cache locations
# Layout of the out_cache_loc
# [ topk 0 ] [ topk 1 ]
# [iter=0, iter=1, iter=2] [iter=0, iter=1, iter=2]
if self.page_size == 1:
out_cache_loc, token_to_kv_pool_state_backup = alloc_token_slots(
batch.tree_cache,
num_seqs * self.speculative_num_steps * self.topk,
backup_state=True,
)
duplicate_cache_len = 0
source_cache_loc, target_cache_loc, last_page_lens_cumsum = None, None, None
else:
if self.topk == 1:
prefix_lens, seq_lens, last_loc = get_last_loc_large_page_size_top_k_1(
batch.req_to_token_pool.req_to_token,
batch.req_pool_indices,
batch.seq_lens,
self.speculative_num_steps,
)
prefix_lens_cpu = batch.seq_lens_cpu
seq_lens_cpu = batch.seq_lens_cpu + self.speculative_num_steps
extend_num_tokens = num_seqs * self.speculative_num_steps
duplicate_cache_len = 0
source_cache_loc, target_cache_loc, last_page_lens_cumsum = (
None,
None,
None,
)
else:
# In this case, the last partial page needs to be duplicated.
# KV cache layout in batch.req_to_token_pool.req_to_token:
#
# | -------- | -- xxxx .. | -- xxxx .. | -- xxxx .. |
# prefix top-k = 0 tok-k = 1 top-k = 2
#
# "-" means prefix tokens
# "x" means speculative draft tokens
# "." means padded tokens
# TODO(lmzheng): The current implementation is still a fake support
# for page size > 1. In the `assign_draft_cache_locs` below,
# we directly move the indices instead of the real kv cache.
# This only works when the kernel backend runs with page size = 1.
# If the kernel backend runs with page size > 1, we need to
# duplicate the real KV cache. The overhead of duplicating KV
# cache seems okay because the draft KV cache only has one layer.
# see a related copy operation in MHATokenToKVPool::move_kv_cache.
(
prefix_lens,
seq_lens,
last_loc,
self.num_new_pages_per_topk,
self.extend_lens,
_,
) = get_last_loc_large_page_size_large_top_k(
batch.req_to_token_pool.req_to_token,
batch.req_pool_indices,
batch.seq_lens,
self.speculative_num_steps,
self.topk,
self.page_size,
)
prefix_lens_cpu = batch.seq_lens_cpu
last_page_lens = prefix_lens_cpu % self.page_size
num_new_pages_per_topk = (
last_page_lens + self.speculative_num_steps + self.page_size - 1
) // self.page_size
seq_lens_cpu = (
prefix_lens_cpu // self.page_size * self.page_size
+ num_new_pages_per_topk * (self.page_size * self.topk)
)
extend_num_tokens = torch.sum((seq_lens_cpu - prefix_lens_cpu)).item()
out_cache_loc, token_to_kv_pool_state_backup = (
alloc_paged_token_slots_extend(
batch.tree_cache,
prefix_lens,
prefix_lens_cpu,
seq_lens,
seq_lens_cpu,
last_loc,
extend_num_tokens,
backup_state=True,
)
)
last_page_lens_cumsum = torch.cumsum(last_page_lens, dim=0)
duplicate_cache_len = torch.sum(last_page_lens).item() * (self.topk - 1)
target_cache_loc = torch.zeros(
duplicate_cache_len, dtype=torch.int32, device=self.device
)
source_cache_loc = torch.zeros(
duplicate_cache_len, dtype=torch.int32, device=self.device
)
assign_draft_cache_locs[(num_seqs,)](
batch.req_pool_indices,
batch.req_to_token_pool.req_to_token,
batch.seq_lens,
self.extend_lens,
self.num_new_pages_per_topk,
out_cache_loc,
source_cache_loc,
target_cache_loc,
last_page_lens_cumsum,
duplicate_cache_len,
batch.req_to_token_pool.req_to_token.shape[1],
self.topk,
self.speculative_num_steps,
self.page_size,
next_power_of_2(num_seqs),
next_power_of_2(self.speculative_num_steps),
)
if self.page_size > 1 and self.topk > 1:
# Remove padded slots
out_cache_loc = out_cache_loc[
: num_seqs * self.topk * self.speculative_num_steps
]
batch.out_cache_loc = out_cache_loc
batch.seq_lens_sum = torch.sum(batch.seq_lens).item()
batch.return_hidden_states = False
spec_info.positions = batch.seq_lens.repeat_interleave(self.topk, dim=0)
self.token_to_kv_pool_allocator.restore_state(token_to_kv_pool_state_backup)
def _draft_preprocess_idle(self, batch: ScheduleBatch):
batch.spec_info = EagleDraftInput.create_idle_input(
device=self.device,
hidden_size=self.model_config.hidden_size,
dtype=self.model_config.dtype,
topk=self.topk * self.speculative_num_steps,
capture_hidden_mode=CaptureHiddenMode.LAST,
)
def draft(self, batch: ScheduleBatch):
# Parse args
if batch.forward_mode.is_idle():
self._draft_preprocess_idle(batch)
else:
self._draft_preprocess_decode(batch)
spec_info = batch.spec_info
assert isinstance(spec_info, EagleDraftInput)
spec_info.capture_hidden_mode = CaptureHiddenMode.LAST
spec_info.num_tokens_per_batch = self.topk
spec_info.num_tokens_for_logprob_per_batch = self.topk
batch.return_hidden_states = False
# Get forward batch
model_worker_batch = batch.get_model_worker_batch()
assert model_worker_batch.capture_hidden_mode == CaptureHiddenMode.LAST
forward_batch = ForwardBatch.init_new(
model_worker_batch, self.mtp_model_runner(0)
)
forward_batch.can_run_dp_cuda_graph = False
forward_batch.return_hidden_states_before_norm = True
# Parse args
assert isinstance(spec_info, EagleDraftInput)
topk_p, topk_index, hidden_states = (
spec_info.topk_p,
spec_info.topk_index,
spec_info.hidden_states,
)
# Return values
score_list: List[torch.Tensor] = []
token_list: List[torch.Tensor] = []
parents_list: List[torch.Tensor] = []
# Forward multiple steps
scores = None
input_ids, hidden_states, scores, tree_info = select_top_k_tokens(
0, topk_p, topk_index, hidden_states, scores, self.topk
)
if self.speculative_num_steps == 1:
score_list.append(tree_info[0])
token_list.append(tree_info[1])
parents_list.append(tree_info[2])
else:
for i in range(self.speculative_num_steps):
score_list.append(tree_info[0][:, :, i].unsqueeze(-1))
token_index = tree_info[1][:, i].unsqueeze(-1)
token_list.append(token_index)
if i == 0:
parents_list.append(tree_info[2])
else:
parents_list.append(
torch.full(
(tree_info[2].size(0), 1),
i,
dtype=torch.long,
device=self.device,
)
)
parent_list, top_scores_index, draft_tokens = organize_draft_results(
score_list, token_list, parents_list, self.speculative_num_draft_tokens
)
if batch.forward_mode.is_idle():
return EagleVerifyInput.create_idle_input(
self.topk,
self.speculative_num_steps,
self.speculative_num_draft_tokens,
)
(
tree_mask,
position,
retrive_index,
retrive_next_token,
retrive_next_sibling,
draft_tokens,
) = build_tree_kernel_efficient(
spec_info.verified_id,
parent_list,
top_scores_index,
draft_tokens,
batch.seq_lens,
batch.seq_lens_sum,
self.topk,
self.speculative_num_steps,
self.speculative_num_draft_tokens,
)
return EagleVerifyInput(
draft_token=draft_tokens,
custom_mask=tree_mask,
positions=position,
retrive_index=retrive_index,
retrive_next_token=retrive_next_token,
retrive_next_sibling=retrive_next_sibling,
retrive_cum_len=None,
spec_steps=self.speculative_num_steps,
topk=self.topk,
draft_token_num=self.server_args.speculative_num_draft_tokens,
capture_hidden_mode=CaptureHiddenMode.FULL,
seq_lens_sum=forward_batch.seq_lens_sum,
seq_lens_cpu=forward_batch.seq_lens_cpu,
)
def clear_cache_pool(self):
# allocator and kv cache pool are shared with target worker
pass
def verify(self, batch: ScheduleBatch, spec_info: EagleVerifyInput):
spec_info.prepare_for_verify(batch, self.page_size)
batch.return_hidden_states = False
batch.forward_mode = (
ForwardMode.TARGET_VERIFY
if not batch.forward_mode.is_idle()
else ForwardMode.IDLE
)
batch.spec_info = spec_info
model_worker_batch = batch.get_model_worker_batch(
seq_lens_cpu_cache=spec_info.seq_lens_cpu
)
assert model_worker_batch.capture_hidden_mode == spec_info.capture_hidden_mode
model_worker_batch.return_hidden_states_before_norm = True
if batch.has_grammar:
retrieve_next_token_cpu = spec_info.retrive_next_token.cpu()
retrieve_next_sibling_cpu = spec_info.retrive_next_sibling.cpu()
draft_tokens_cpu = spec_info.draft_token.view(
spec_info.retrive_next_token.shape
).cpu()
# Forward
batch_result = self.target_worker.forward_batch_generation(
model_worker_batch, is_verify=True
)
logits_output, can_run_cuda_graph = (
batch_result.logits_output,
batch_result.can_run_cuda_graph,
)
vocab_mask = None
if batch.has_grammar:
# Generate the logit mask for structured output.
# Overlap the CPU operations for bitmask generation with the forward pass.
vocab_mask = generate_token_bitmask(
batch.reqs,
spec_info,
retrieve_next_token_cpu,
retrieve_next_sibling_cpu,
draft_tokens_cpu,
batch.sampling_info.vocab_size,
)
if vocab_mask is not None:
assert spec_info.grammar is not None
vocab_mask = vocab_mask.to(spec_info.retrive_next_token.device)
# NOTE (sk): otherwise, this vocab mask will be the one from the previous extend stage
# and will be applied to produce wrong results
batch.sampling_info.vocab_mask = None
if self.enable_nan_detection:
detect_nan(logits_output)
spec_info.hidden_states = logits_output.hidden_states
res: EagleVerifyOutput = spec_info.verify(
batch,
logits_output,
self.token_to_kv_pool_allocator,
self.page_size,
vocab_mask,
)
# Post process based on verified outputs.
# Pick indices that we care (accepted)
logits_output.next_token_logits = logits_output.next_token_logits[
res.accepted_indices
]
logits_output.hidden_states = logits_output.hidden_states[res.accepted_indices]
if self.target_worker.model_runner.hybrid_gdn_config is not None:
accepted_length = (
torch.tensor(
res.accept_length_per_req_cpu,
device=logits_output.hidden_states.device,
dtype=torch.int64,
)
+ 1
)
# If topk > 1, we need to use retrieve_next_token and retrieve_next_sibling to handle the eagle tree custom attention mask
# res.accepted_indices.shape[0] > 0 skips DP attn idle batch
if spec_info.topk > 1 and res.accepted_indices.shape[0] > 0:
# accepted_indices=[0,2,3,4,5,7,9,10,11], accepted_length=[4, 3, 2], cumulative_accepted_lengths=[4, 7, 9]
# first_token_indices_per_req=prepend(0, accepted_indices[cumulative_accepted_lengths[:-1]]) = [0, 5, 10]
# last_token_indices_per_req=accepted_indices[cumulative_accepted_lengths - 1] = [4, 9, 11] (last token ID of each req)
# max_relative_indices_per_req = [4,4,1]; those are the per-req spec-decoding step offsets that contain the correct mamba caches
cumulative_accepted_lengths = torch.cumsum(accepted_length, dim=0)
req_start_positions = torch.cat(
[
torch.zeros(
1,
dtype=cumulative_accepted_lengths.dtype,
device=cumulative_accepted_lengths.device,
),
cumulative_accepted_lengths[:-1],
]
)
first_token_indices_per_req = res.accepted_indices[req_start_positions]
last_token_indices_per_req = res.accepted_indices[
cumulative_accepted_lengths - 1
]
max_relative_indices_per_req = (
last_token_indices_per_req - first_token_indices_per_req
)
else:
max_relative_indices_per_req = accepted_length - 1
self.target_worker.model_runner.attn_backend.update_mamba_state_after_mtp_verify(
max_relative_indices_per_req, self.target_worker.model_runner.model
)
if batch.return_logprob:
self.add_logprob_values(batch, res, logits_output)
# Prepare the batch for the next draft forwards.
batch.forward_mode = (
ForwardMode.DECODE if not batch.forward_mode.is_idle() else ForwardMode.IDLE
)
batch.spec_info = res.draft_input
return logits_output, res, model_worker_batch, can_run_cuda_graph
def add_logprob_values(
self,
batch: ScheduleBatch,
res: EagleVerifyOutput,
logits_output: LogitsProcessorOutput,
):
# Extract args
logits_output = res.logits_output
top_logprobs_nums = batch.top_logprobs_nums
token_ids_logprobs = batch.token_ids_logprobs
accepted_indices = res.accepted_indices
assert len(accepted_indices) == len(logits_output.next_token_logits)
temperatures = batch.sampling_info.temperatures
num_draft_tokens = batch.spec_info.draft_token_num
# acceptance indices are the indices in a "flattened" batch.
# dividing it to num_draft_tokens will yield the actual batch index.
temperatures = temperatures[accepted_indices // num_draft_tokens]
if SGLANG_RETURN_ORIGINAL_LOGPROB:
logprobs = torch.nn.functional.log_softmax(
logits_output.next_token_logits, dim=-1
)
else:
logprobs = torch.nn.functional.log_softmax(
logits_output.next_token_logits / temperatures, dim=-1
)
batch_next_token_ids = res.verified_id
num_tokens_per_req = [accept + 1 for accept in res.accept_length_per_req_cpu]
# We should repeat top_logprobs_nums to match num_tokens_per_req.
top_logprobs_nums_repeat_interleaved = []
token_ids_logprobs_repeat_interleaved = []
for num, num_tokens in zip(top_logprobs_nums, num_tokens_per_req):
top_logprobs_nums_repeat_interleaved.extend([num] * num_tokens)
for token_ids, num_tokens in zip(token_ids_logprobs, num_tokens_per_req):
token_ids_logprobs_repeat_interleaved.extend([token_ids] * num_tokens)
# Extract logprobs
if any(x > 0 for x in top_logprobs_nums):
(
logits_output.next_token_top_logprobs_val,
logits_output.next_token_top_logprobs_idx,
) = get_top_logprobs(
logprobs,
top_logprobs_nums_repeat_interleaved,
)
if any(x is not None for x in token_ids_logprobs):
(
logits_output.next_token_token_ids_logprobs_val,
logits_output.next_token_token_ids_logprobs_idx,
) = get_token_ids_logprobs(
logprobs,
token_ids_logprobs_repeat_interleaved,
)
logits_output.next_token_logprobs = logprobs[
torch.arange(len(batch_next_token_ids), device=batch.sampling_info.device),
batch_next_token_ids,
]
# Add output logprobs to the request
pt = 0
next_token_logprobs = logits_output.next_token_logprobs.tolist()
verified_ids = batch_next_token_ids.tolist()
for req, num_tokens in zip(batch.reqs, num_tokens_per_req, strict=True):
for _ in range(num_tokens):
if req.return_logprob:
req.output_token_logprobs_val.append(next_token_logprobs[pt])
req.output_token_logprobs_idx.append(verified_ids[pt])
if req.top_logprobs_num > 0:
req.output_top_logprobs_val.append(
res.logits_output.next_token_top_logprobs_val[pt]
)
req.output_top_logprobs_idx.append(
res.logits_output.next_token_top_logprobs_idx[pt]
)
pt += 1
def forward_draft_extend(
self,
batch: ScheduleBatch,
hidden_states: torch.Tensor,
next_token_ids: torch.Tensor,
seq_lens_cpu: Optional[torch.Tensor],
):
"""Run draft model extend. This API modifies the states of the batch.
Args:
batch: The batch to run.
hidden_states: Hidden states from the target model forward
next_token_ids: Next token ids generated from the target forward.
"""
batch.spec_info = EagleDraftInput(
hidden_states=hidden_states,
verified_id=next_token_ids,
num_tokens_per_batch=1,
num_tokens_for_logprob_per_batch=1,
)
batch.return_hidden_states = False
batch.spec_info.prepare_for_extend(batch)
batch.spec_info.capture_hidden_mode = CaptureHiddenMode.LAST
model_worker_batch = batch.get_model_worker_batch(
seq_lens_cpu_cache=seq_lens_cpu
)
forward_batch = ForwardBatch.init_new(
model_worker_batch, self.mtp_model_runner(0)
)
forward_batch.return_logprob = False
forward_batch.return_hidden_states_before_norm = True
topk_p_list = []
topk_index_list = []
for step in range(self.speculative_num_steps):
logits_output, _ = self.mtp_model_runner(step).forward(forward_batch)
if self.enable_nan_detection:
detect_nan(logits_output)
probs = torch.softmax(logits_output.next_token_logits, dim=-1)
topk_p, topk_index = fast_topk(probs, self.topk, dim=-1)
topk_p_list.append(topk_p)
topk_index_list.append(topk_index)
pt = 0
if forward_batch.extend_seq_lens is not None:
for i, extend_len in enumerate(forward_batch.extend_seq_lens):
input_ids = forward_batch.input_ids[pt : pt + extend_len]
forward_batch.input_ids[pt : pt + extend_len] = torch.cat(
(input_ids[1:], topk_index[i].reshape(1))
)
pt += extend_len
assert isinstance(forward_batch.spec_info, EagleDraftInput)
assert forward_batch.spec_info is batch.spec_info
forward_batch.spec_info.topk_p = torch.cat(topk_p_list, dim=1)
forward_batch.spec_info.topk_index = torch.cat(topk_index_list, dim=1)
has_finished, unfinished_req_index = False, []
for i, req in enumerate(batch.reqs):
if req.finished():
has_finished = True
else:
unfinished_req_index.append(i)
if has_finished:
unfinished_index_device = torch.tensor(
unfinished_req_index,
dtype=torch.int64,
device=batch.spec_info.topk_p.device,
)
batch.spec_info.filter_batch(
unfinished_index_device, has_been_filtered=False
)
def forward_draft_extend_after_decode(self, batch: ScheduleBatch):
assert isinstance(batch.spec_info, EagleDraftInput)
# Backup fields that will be modified in-place
seq_lens_backup = batch.seq_lens.clone()
seq_lens_cpu_backup = batch.seq_lens_cpu.clone()
req_pool_indices_backup = batch.req_pool_indices
accept_length_backup = batch.spec_info.accept_length
return_logprob_backup = batch.return_logprob
input_is_idle = batch.forward_mode.is_idle()
if not input_is_idle and batch.spec_info.verified_id.numel() == 0:
batch = batch.copy()
batch.prepare_for_idle()
hidden_size = (
self.model_config.hidden_size * 3
if self.speculative_algorithm.is_eagle3()
else self.model_config.hidden_size
)
batch.spec_info = EagleDraftInput.create_idle_input(
device=self.device,
hidden_size=hidden_size,
dtype=self.model_config.dtype,
topk=self.topk,
capture_hidden_mode=CaptureHiddenMode.LAST,
)
batch.spec_info.num_tokens_per_batch = self.speculative_num_steps + 1
batch.spec_info.num_tokens_for_logprob_per_batch = 1
batch.spec_info.prepare_extend_after_decode(
batch,
self.speculative_num_steps,
)
batch.forward_mode = (
ForwardMode.DRAFT_EXTEND
if not batch.forward_mode.is_idle()
else ForwardMode.IDLE
)
batch.return_hidden_states = False
model_worker_batch = batch.get_model_worker_batch()
assert model_worker_batch.capture_hidden_mode == CaptureHiddenMode.LAST
forward_batch = ForwardBatch.init_new(
model_worker_batch, self.mtp_model_runner(0)
)
forward_batch.return_hidden_states_before_norm = True
if forward_batch.seq_lens_cpu is not None:
forward_batch.seq_lens_sum = forward_batch.seq_lens_cpu.sum().item()
else:
forward_batch.seq_lens_sum = batch.seq_lens.sum().item()
topk_p_list = []
topk_index_list = []
# Run
for step in range(self.speculative_num_steps):
can_cuda_graph = len(
self.cuda_graph_runner_for_draft_extend_list
) and self.cuda_graph_runner_for_draft_extend_list[step].can_run(
forward_batch
)
if can_cuda_graph:
logits_output = self.cuda_graph_runner_for_draft_extend_list[
step
].replay(forward_batch)
else:
forward_batch.can_run_dp_cuda_graph = False
if not forward_batch.forward_mode.is_idle():
self.mtp_model_runner(step).attn_backend.init_forward_metadata(
forward_batch
)
logits_output, _ = self.mtp_model_runner(step).forward(
forward_batch, skip_attn_backend_init=True
)
if self.enable_nan_detection:
detect_nan(logits_output)
probs = torch.softmax(logits_output.next_token_logits, dim=-1)
topk_p, topk_index = fast_topk(probs, self.topk, dim=-1)
topk_p_list.append(topk_p)
topk_index_list.append(topk_index)
pt = 0
if forward_batch.extend_seq_lens is not None:
for i, extend_len in enumerate(forward_batch.extend_seq_lens):
input_ids = forward_batch.input_ids[pt : pt + extend_len]
forward_batch.input_ids[pt : pt + extend_len] = torch.cat(
(input_ids[1:], topk_index[i].reshape(1))
)
pt += extend_len
forward_batch.spec_info.topk_p = torch.cat(topk_p_list, dim=1)
forward_batch.spec_info.topk_index = torch.cat(topk_index_list, dim=1)
# Restore backup.
# This is because `seq_lens` can be modified in `prepare_extend_after_decode`
batch.forward_mode = (
ForwardMode.DECODE if not input_is_idle else ForwardMode.IDLE
)
batch.seq_lens = seq_lens_backup
batch.seq_lens_cpu = seq_lens_cpu_backup
batch.req_pool_indices = req_pool_indices_backup
batch.spec_info.accept_length = accept_length_backup
batch.return_logprob = return_logprob_backup
@@ -0,0 +1,750 @@
# Copyright 2023-2024 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
import contextlib
import logging
from typing import List, Optional, Tuple
import torch
from sglang.srt.environ import envs
from sglang.srt.layers.moe.utils import speculative_moe_backend_context
from sglang.srt.managers.schedule_batch import ModelWorkerBatch
from sglang.srt.managers.scheduler import GenerationBatchResult
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode, ForwardBatch
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.base_spec_worker import BaseDraftWorker, BaseSpecWorker
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
from sglang.srt.speculative.eagle_info_v2 import (
assign_extend_cache_locs,
fill_accepted_out_cache_loc,
fill_new_verified_id,
)
from sglang.srt.speculative.eagle_utils import TreeMaskMode, build_tree_kernel_efficient
from sglang.srt.speculative.mtp_draft_extend_cuda_graph_runner import (
MTPMultiStepDraftExtendCudaGraphRunner,
)
from sglang.srt.speculative.mtp_utils import (
assign_hidden_states_pool_triton,
rotate_input_ids_triton,
)
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import (
detect_nan,
draft_tp_context,
select_top_k_tokens,
)
from sglang.srt.utils.common import empty_context, fast_topk, next_power_of_2
logger = logging.getLogger(__name__)
def _get_plan_stream(
device: str,
) -> Tuple[any, contextlib.AbstractContextManager]:
if envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.get():
plan_stream = torch.get_device_module(device).Stream()
plan_stream_ctx = torch.get_device_module(device).stream(plan_stream)
return plan_stream, plan_stream_ctx
else:
return None, contextlib.nullcontext()
class MTPDraftWorker(BaseDraftWorker):
def __init__(
self,
server_args: ServerArgs,
gpu_id: int,
tp_rank: int,
dp_rank: int,
moe_ep_rank: int,
nccl_port: int,
target_worker: TpModelWorker,
):
# copy args
self.server_args = server_args
self.gpu_id = gpu_id
self.tp_rank = tp_rank
self.dp_rank = dp_rank
self.moe_ep_rank = moe_ep_rank
self.nccl_port = nccl_port
self.target_worker = target_worker
self.draft_extend_attn_backend_list = []
self.model_config = target_worker.model_config
# Args for easy access
self.device = server_args.device
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm
)
# Set constant
EagleDraftInput.ALLOC_LEN_PER_DECODE = max(
self.speculative_num_steps * self.topk, self.speculative_num_draft_tokens
)
# Do not capture cuda graph in `TpModelWorker` init,
# will capture later with init_cuda_graphs()
backup_disable_cuda_graph = server_args.disable_cuda_graph
server_args.disable_cuda_graph = True
# Share the allocator with a target worker.
# Draft and target worker own their own KV cache pools.
self.req_to_token_pool, self.token_to_kv_pool_allocator = (
target_worker.get_memory_pool()
)
with empty_context(), speculative_moe_backend_context():
# Init draft worker
self.draft_worker = TpModelWorker(
server_args=server_args,
gpu_id=gpu_id,
tp_rank=tp_rank,
pp_rank=0, # FIXME
dp_rank=dp_rank,
moe_ep_rank=moe_ep_rank,
nccl_port=nccl_port,
is_draft_worker=True,
req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
is_mtp_worker=True,
)
# Alias for better readability
# self.draft_runner = self.draft_worker.model_runner
self.draft_runner_list = self.draft_worker.model_runner_list
self.init_lm_head()
# Used for KV Cache reversion
self.req_to_hidden_states_pool = torch.empty(
(
self.req_to_token_pool.size,
self.speculative_num_steps - 1,
self.model_config.hidden_size,
),
dtype=self.model_config.dtype,
device=self.device,
)
# Init attention backend and cuda graphs
for i in range(self.speculative_num_steps):
self.draft_runner_list[i].server_args.disable_cuda_graph = (
backup_disable_cuda_graph
)
self.draft_tp_context = (
draft_tp_context if server_args.enable_dp_attention else empty_context
)
with self.draft_tp_context(
self.draft_runner_list[0].tp_group
), speculative_moe_backend_context():
self.init_attention_backend()
self.init_cuda_graphs()
self.tree_mask_mode = TreeMaskMode.FULL_MASK
self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device)
def mtp_model_runner(self, step: int):
return self.draft_runner_list[step]
def init_lm_head(self):
embed, head = self.target_worker.model_runner.model.get_embed_and_head()
# Share the embedding and lm_head
for i in range(self.speculative_num_steps):
self.draft_runner_list[i].model.set_embed_and_head(embed, head)
def init_attention_backend(self):
# Create attn backends
self.draft_extend_attn_backend_list = []
for step in range(self.speculative_num_steps):
from sglang.srt.layers.attention.flashattention_backend import (
FlashAttentionBackend,
)
self.draft_extend_attn_backend_list.append(
FlashAttentionBackend(
model_runner=self.draft_runner_list[step],
skip_prefill=False,
speculative_step_id=step,
)
)
self.draft_runner_list[step].attn_backend = (
self.draft_extend_attn_backend_list[-1]
)
def init_cuda_graphs(self):
"""Capture cuda graphs."""
self.cuda_graph_runner = None
self.cuda_graph_runner_for_draft_extend = None
if self.server_args.disable_cuda_graph:
return
self.cuda_graph_runner_for_draft_extend = (
MTPMultiStepDraftExtendCudaGraphRunner(self)
)
def reset_cuda_graph_buffers(self, forward_batch, batch_result):
if self.cuda_graph_runner_for_draft_extend:
self.cuda_graph_runner_for_draft_extend.reset_buffers(
forward_batch, batch_result
)
def draft(self, model_worker_batch: ModelWorkerBatch):
draft_input: EagleDraftInput = model_worker_batch.spec_info
forward_batch, can_cuda_graph = draft_input.prepare_for_v2_draft(
self.req_to_token_pool,
model_worker_batch,
self.cuda_graph_runner,
self.draft_runner_list[0],
self.topk,
self.speculative_num_steps,
)
# Run draft
parent_list, top_scores_index, draft_tokens = self.draft_forward(forward_batch)
if model_worker_batch.forward_mode.is_idle():
return EagleVerifyInput.create_idle_input(
self.topk,
self.speculative_num_steps,
self.speculative_num_draft_tokens,
)
# Build tree mask
# Directly write to cuda graph buffers for verify attn
tree_mask_buf, position_buf = (
self.target_worker.model_runner.attn_backend.get_verify_buffers_to_fill_after_draft()
)
(
tree_mask,
position,
retrive_index,
retrive_next_token,
retrive_next_sibling,
draft_tokens,
) = build_tree_kernel_efficient(
draft_input.verified_id,
parent_list,
top_scores_index,
draft_tokens,
model_worker_batch.seq_lens,
model_worker_batch.seq_lens_sum,
self.topk,
self.speculative_num_steps,
self.speculative_num_draft_tokens,
self.tree_mask_mode,
tree_mask_buf,
position_buf,
)
return EagleVerifyInput(
draft_token=draft_tokens,
custom_mask=tree_mask,
positions=position,
retrive_index=retrive_index,
retrive_next_token=retrive_next_token,
retrive_next_sibling=retrive_next_sibling,
retrive_cum_len=None,
spec_steps=self.speculative_num_steps,
topk=self.topk,
draft_token_num=self.speculative_num_draft_tokens,
capture_hidden_mode=None,
seq_lens_sum=None,
seq_lens_cpu=None,
)
def draft_forward(self, forward_batch: ForwardBatch):
# Parse args
spec_info: EagleDraftInput = forward_batch.spec_info
topk_p, topk_index, hidden_states = (
spec_info.topk_p,
spec_info.topk_index,
spec_info.hidden_states,
)
# Return values
score_list: List[torch.Tensor] = []
token_list: List[torch.Tensor] = []
parents_list: List[torch.Tensor] = []
# Forward multiple steps
scores = None
_, hidden_states, scores, tree_info = select_top_k_tokens(
0, topk_p, topk_index, hidden_states, scores, self.topk
)
if self.speculative_num_steps == 1:
score_list.append(tree_info[0])
token_list.append(tree_info[1])
parents_list.append(tree_info[2])
else:
for i in range(self.speculative_num_steps):
score_list.append(tree_info[0][:, :, i].unsqueeze(-1))
token_index = tree_info[1][:, i].unsqueeze(-1)
token_list.append(token_index)
if i == 0:
parents_list.append(tree_info[2])
else:
parents_list.append(
torch.full(
(tree_info[2].size(0), 1),
i,
dtype=torch.long,
device="cuda",
)
)
# Organize the results
score_list = torch.cat(score_list, dim=1).flatten(
1
) # b, n, topk; n= 1 + (num_steps-1) * self.topk
ss_token_list = torch.cat(
token_list, dim=1
) # b, (self.topk + (num_steps-1) * self.topk)
top_scores = torch.topk(
score_list, self.speculative_num_draft_tokens - 1, dim=-1
)
top_scores_index = top_scores.indices
top_scores_index = torch.sort(top_scores_index).values
draft_tokens = torch.gather(ss_token_list, index=top_scores_index, dim=1)
if len(parents_list) > 1:
parent_list = torch.cat(parents_list[:-1], dim=1)
else:
batch_size = parents_list[0].shape[0]
parent_list = torch.empty(batch_size, 0, device=parents_list[0].device)
return parent_list, top_scores_index, draft_tokens
def draft_extend(self):
pass
def _draft_extend_for_prefill(
self,
batch: ModelWorkerBatch,
target_hidden_states: torch.Tensor,
next_token_ids: torch.Tensor,
):
"""
Run draft model extend to correctly fill the KV cache.
Args:
batch: The batch to run.
target_hidden_states: Hidden states from the target model forward
next_token_ids: Next token ids generated from the target forward.
"""
# Construct spec_info
next_draft_input = EagleDraftInput(
hidden_states=target_hidden_states,
verified_id=next_token_ids,
new_seq_lens=batch.seq_lens,
# draft mode is same with decode mode, only 1 num token per batch
num_tokens_per_batch=1,
num_tokens_for_logprob_per_batch=1,
)
batch.spec_info = next_draft_input
# Run forward
forward_batch = ForwardBatch.init_new(batch, self.draft_runner_list[0])
forward_batch.return_hidden_states_before_norm = True
# Construct input_ids
if not batch.forward_mode.is_idle():
rotate_input_ids_triton(
forward_batch.input_ids,
forward_batch.extend_start_loc,
forward_batch.extend_seq_lens,
next_token_ids,
)
topk_p_list = []
topk_index_list = []
for step in range(self.speculative_num_steps):
logits_output, _ = self.draft_runner_list[step].forward(forward_batch)
probs = torch.softmax(logits_output.next_token_logits, dim=-1)
topk_p, topk_index = fast_topk(probs, self.topk, dim=-1)
topk_p_list.append(topk_p)
topk_index_list.append(topk_index)
if forward_batch.extend_seq_lens is not None:
rotate_input_ids_triton(
forward_batch.input_ids,
forward_batch.extend_start_loc,
forward_batch.extend_seq_lens,
topk_index,
)
next_draft_input.topk_p = torch.cat(topk_p_list, dim=1)
next_draft_input.topk_index = torch.cat(topk_index_list, dim=1)
# next_draft_input.hidden_states = logits_output.hidden_states
# Update req_to_hidden_states_pool for KV Cache reversion
if forward_batch.extend_seq_lens is not None:
assign_hidden_states_pool_triton(
target_hidden_states,
forward_batch.req_pool_indices,
self.req_to_hidden_states_pool,
self.speculative_num_steps - 1,
forward_batch.batch_size,
forward_batch.extend_seq_lens,
forward_batch.extend_start_loc,
)
return next_draft_input
def _draft_extend_for_decode(
self, batch: ModelWorkerBatch, batch_result: GenerationBatchResult
):
# Batch 2: Draft extend
draft_input = EagleDraftInput(
hidden_states=batch_result.logits_output.hidden_states,
num_tokens_per_batch=self.speculative_num_steps + 1,
num_tokens_for_logprob_per_batch=1,
)
# Prepare for draft extend in a separate stream
# Notice that here we use batch_result.next_token_ids as the input ids
with self.plan_stream_ctx:
forward_batch = draft_input.prepare_for_extend_to_fill_draft_kvcache(
batch,
batch_result.next_token_ids,
self.speculative_num_draft_tokens,
self.draft_runner_list[0],
self.cuda_graph_runner_for_draft_extend,
)
forward_batch.return_hidden_states_before_norm = True
if self.plan_stream:
torch.get_device_module(self.device).current_stream().wait_stream(
self.plan_stream
)
# Run draft extend batch in the main compute stream
can_cuda_graph = (
self.cuda_graph_runner_for_draft_extend
and self.cuda_graph_runner_for_draft_extend.can_run(forward_batch)
)
ret_topk_p_list = []
ret_topk_index_list = []
next_token_ids_backup = batch_result.next_token_ids.clone()
if can_cuda_graph:
self.reset_cuda_graph_buffers(forward_batch, batch_result)
else:
logger.warning_once(
f"can't use cuda graph for draft extend! may have correctness issue!"
)
select_index = (
torch.arange(len(batch.seq_lens), device=self.device)
* self.speculative_num_draft_tokens
+ batch_result.accept_lens
- 1
)
for step in range(self.speculative_num_steps):
# log_info_on_rank0(logger, f"step: {step}, forward_batch.input_ids: {forward_batch.input_ids}")
if can_cuda_graph:
draft_logits_output = (
self.cuda_graph_runner_for_draft_extend.get_runner(step).replay(
forward_batch, init_state=(step == 0)
)
)
ret_topk_p, ret_topk_index = (
draft_logits_output.topk_p,
draft_logits_output.topk_index,
)
else:
draft_logits_output, _ = self.draft_runner_list[step].forward(
forward_batch, skip_attn_backend_init=True
)
probs = torch.softmax(
draft_logits_output.next_token_logits[select_index], dim=-1
)
ret_topk_p, ret_topk_index = fast_topk(probs, self.topk, dim=-1)
if forward_batch.extend_seq_lens is not None:
rotate_input_ids_triton(
forward_batch.input_ids,
forward_batch.extend_start_loc,
forward_batch.extend_seq_lens,
ret_topk_index,
select_index,
)
ret_topk_p_list.append(ret_topk_p)
ret_topk_index_list.append(ret_topk_index)
# Update req_to_hidden_states_pool for KV Cache reversion
if (
self.cuda_graph_runner_for_draft_extend is not None
and forward_batch.extend_seq_lens is not None
):
last_cuda_graph_runner = (
self.cuda_graph_runner_for_draft_extend.get_last_runner()
)
assign_hidden_states_pool_triton(
last_cuda_graph_runner.hidden_states,
last_cuda_graph_runner.req_pool_indices,
self.req_to_hidden_states_pool,
self.speculative_num_steps - 1,
forward_batch.batch_size,
last_cuda_graph_runner.extend_seq_lens,
last_cuda_graph_runner.extend_start_loc,
)
# Reorganize the spec info for the next batch
# draft_logits_output.next_token_logits = draft_logits_output.next_token_logits[
# select_index
# ]
# draft_logits_output.hidden_states = draft_logits_output.hidden_states[
# select_index
# ]
batch_result.next_token_ids = next_token_ids_backup
# Construct the return values
next_draft_input = batch_result.next_draft_input
(
next_draft_input.topk_p,
next_draft_input.topk_index,
next_draft_input.hidden_states,
) = (
torch.cat(ret_topk_p_list, dim=1).clone(),
torch.cat(ret_topk_index_list, dim=1).clone(),
None,
)
class MTPWorkerV2(BaseSpecWorker):
def __init__(
self,
server_args: ServerArgs,
gpu_id: int,
tp_rank: int,
dp_rank: Optional[int],
moe_ep_rank: int,
nccl_port: int,
target_worker: TpModelWorker,
):
# Parse arguments
self.server_args = server_args
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.enable_nan_detection = server_args.enable_nan_detection
self.gpu_id = gpu_id
self.device = server_args.device
self._target_worker = target_worker
self.page_size = server_args.page_size
self.speculative_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm
)
self.req_to_token_pool, self.token_to_kv_pool_allocator = (
target_worker.get_memory_pool()
)
# Override the context length of the draft model to be the same as the target model.
server_args.context_length = target_worker.model_runner.model_config.context_len
self._draft_worker = MTPDraftWorker(
server_args, gpu_id, tp_rank, dp_rank, moe_ep_rank, nccl_port, target_worker
)
# Some dummy tensors
self.num_new_pages_per_topk = torch.empty(
(), dtype=torch.int64, device=self.device
)
self.extend_lens = torch.empty((), dtype=torch.int64, device=self.device)
self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device)
@property
def target_worker(self):
return self._target_worker
@property
def draft_worker(self):
return self._draft_worker
def clear_cache_pool(self):
# allocator and kv cache pool are shared with target worker, which are cleared in scheduler
pass
def forward_batch_generation(self, model_worker_batch: ModelWorkerBatch):
if (
model_worker_batch.forward_mode.is_extend()
or model_worker_batch.is_extend_in_batch
):
# Target prefill
model_worker_batch.capture_hidden_mode = CaptureHiddenMode.FULL
batch_output = self.target_worker.forward_batch_generation(
model_worker_batch
)
# Draft prefill
model_worker_batch.capture_hidden_mode = CaptureHiddenMode.LAST
batch_output.next_draft_input = self.draft_worker._draft_extend_for_prefill(
model_worker_batch,
batch_output.logits_output.hidden_states,
batch_output.next_token_ids,
)
return batch_output
else:
if model_worker_batch.spec_info is None:
model_worker_batch.spec_info = EagleDraftInput.create_idle_input(
device=self.device,
hidden_size=self.target_worker.model_config.hidden_size,
dtype=self.target_worker.model_config.dtype,
topk=self.topk * self.speculative_num_steps,
capture_hidden_mode=CaptureHiddenMode.LAST,
)
draft_input: EagleDraftInput = model_worker_batch.spec_info
verify_input: EagleVerifyInput = self.draft_worker.draft(model_worker_batch)
assert verify_input.is_verify_input()
model_worker_batch.spec_info = verify_input
batch_output = self.verify(model_worker_batch)
self.draft_worker._draft_extend_for_decode(model_worker_batch, batch_output)
return batch_output
def verify(
self,
batch: ModelWorkerBatch,
):
# Since batch.seq_lens is allocated in another stream, we need
# record_stream() to prevent pytorch gc and reuse the gpu memory
# while forward_stream is still running.
batch.seq_lens.record_stream(
torch.get_device_module(self.device).current_stream()
)
# Parse args
verify_input: EagleVerifyInput = batch.spec_info
bs = len(batch.seq_lens)
# Batch 1: Target verify
# Prepare for target verify in a separate stream
with self.plan_stream_ctx:
verify_forward_batch, can_run_cuda_graph = (
verify_input.prepare_for_v2_verify(
self.req_to_token_pool,
batch,
self.target_worker,
)
)
# Correct some buffers due to the overlap plan
if self.plan_stream:
torch.get_device_module(self.device).current_stream().wait_stream(
self.plan_stream
)
# Some values such as custom_mask and position depend on the output of draft,
# so the previous plan step used the wrong values. Here, we need to run the related
# computation again to update them to the correct values.
self.target_worker.model_runner.attn_backend.update_verify_buffers_to_fill_after_draft(
verify_input,
(
self.target_worker.model_runner.graph_runner.bs
if can_run_cuda_graph
else None
),
)
# Run target verify batch in the main compute stream
forward_batch_output = self.target_worker.forward_batch_generation(
model_worker_batch=None,
forward_batch=verify_forward_batch,
is_verify=True,
skip_attn_backend_init=True,
)
logits_output = forward_batch_output.logits_output
# Sample
if self.enable_nan_detection:
detect_nan(logits_output)
(
predict,
accept_length,
accept_index,
) = verify_input.sample(batch, logits_output)
new_seq_lens = batch.seq_lens + accept_length
verify_done = torch.get_device_module(self.device).Event()
verify_done.record()
if not batch.forward_mode.is_idle():
all_verified_id = predict[accept_index]
verified_id = torch.empty_like(accept_length, dtype=torch.int32)
fill_new_verified_id[(bs,)](
all_verified_id,
accept_length,
verified_id,
self.speculative_num_draft_tokens,
)
else:
verified_id = torch.empty((0,), device=self.device, dtype=torch.int32)
# Construct the next draft input
next_draft_input = EagleDraftInput(
verified_id=verified_id,
new_seq_lens=new_seq_lens,
verify_done=verify_done,
)
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=predict,
can_run_cuda_graph=can_run_cuda_graph,
next_draft_input=next_draft_input,
accept_lens=accept_length,
)
def move_accepted_tokens_to_target_kvcache(
self,
batch: ModelWorkerBatch,
accept_index: torch.Tensor,
accept_length: torch.Tensor,
):
"""
Move accepted tokens to the target KV cache.
Args:
batch: The batch to run.
accept_index: The index of the accepted tokens.
accept_length: The length of the accepted tokens.
"""
bs = len(batch.seq_lens)
size = bs * self.speculative_num_draft_tokens
tgt_cache_loc = torch.zeros(
size,
dtype=torch.int64,
device=self.device,
)
accepted_out_cache_loc = torch.zeros(
size, dtype=torch.int64, device=self.device
)
assign_extend_cache_locs[(bs,)](
batch.req_pool_indices,
self.req_to_token_pool.req_to_token,
batch.seq_lens,
batch.seq_lens + accept_length,
tgt_cache_loc,
self.req_to_token_pool.req_to_token.shape[1],
next_power_of_2(bs),
)
fill_accepted_out_cache_loc[(size,)](
accept_index,
batch.out_cache_loc,
accepted_out_cache_loc,
next_power_of_2(size),
)
self.token_to_kv_pool_allocator.get_kvcache().move_kv_cache(
tgt_cache_loc, accepted_out_cache_loc
)
+50 -3
View File
@@ -4,7 +4,7 @@ import logging
import os
import time
from contextlib import contextmanager
from typing import TYPE_CHECKING, List
from typing import TYPE_CHECKING, List, Optional
import torch
import triton
@@ -19,6 +19,8 @@ from sglang.srt.distributed.parallel_state import (
from sglang.srt.environ import envs
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.common import get_last_loc
from sglang.srt.server_args import ServerArgs, get_global_server_args
from sglang.srt.utils import is_cuda, is_hip, is_npu, next_power_of_2
_is_cuda = is_cuda()
@@ -48,6 +50,14 @@ TREE_TRAVERSE_TIME_THRESHOLD = 1 # TODO: set this properly
TREE_SPEC_KERNEL_AVAILABLE = _is_cuda # This kernel is only available for CUDA now
def spec_need_hidden_states(server_args: Optional[ServerArgs] = None) -> bool:
if server_args is None:
server_args = get_global_server_args()
# TODO(lsyin): also skip when 1) step = 1 or 2) standalone draft model
return not server_args.enable_mtp
@triton.jit
def create_extend_after_decode_spec_info(
verified_id,
@@ -465,13 +475,14 @@ def select_top_k_tokens(
if i == 0:
# The first step after extend
input_ids = topk_index.flatten()
hidden_states = hidden_states.repeat_interleave(topk, dim=0)
if hidden_states is not None:
hidden_states = hidden_states.repeat_interleave(topk, dim=0)
scores = topk_p # shape: (b, topk)
tree_info = (
topk_p.unsqueeze(1), # shape: (b, 1, topk)
topk_index, # shape: (b, topk)
torch.arange(-1, topk, dtype=torch.long, device=hidden_states.device)
torch.arange(-1, topk, dtype=torch.long, device=input_ids.device)
.unsqueeze(0)
.repeat(topk_p.shape[0], 1), # shape: (b, topk + 1)
)
@@ -695,3 +706,39 @@ def detect_nan(logits_output: LogitsProcessorOutput):
if torch.any(torch.isnan(logits)):
logger.error("Detected errors during sampling! NaN in the logits.")
raise ValueError("Detected errors during sampling! NaN in the logits.")
# Disable torch.compile for this function because it will be
# even slower.
# @torch.compile(dynamic=True)
def get_last_loc_large_page_size_large_top_k(
req_to_token: torch.Tensor,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
speculative_num_steps: int,
topk: int,
page_size: int,
):
prefix_lens = seq_lens
last_page_lens = prefix_lens % page_size
num_new_pages_per_topk = (
last_page_lens + speculative_num_steps + page_size - 1
) // page_size
seq_lens = prefix_lens // page_size * page_size + num_new_pages_per_topk * (
page_size * topk
)
extend_lens = seq_lens - prefix_lens
last_loc = get_last_loc(
req_to_token,
req_pool_indices,
prefix_lens,
)
return (
prefix_lens,
seq_lens,
last_loc,
num_new_pages_per_topk,
extend_lens,
last_page_lens,
)