feat(model_runner): remove pool/backend refs from ForwardBatch via ForwardContext (#25983)

Co-authored-by: Claude Sonnet 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-05-21 14:01:49 -07:00
committed by GitHub
co-authored by Claude Sonnet 4.6
parent 44ec2ee18d
commit c5251a98a9
77 changed files with 1236 additions and 914 deletions
+62 -7
View File
@@ -1,16 +1,31 @@
from __future__ import annotations from __future__ import annotations
import os import os
from contextlib import contextmanager from contextlib import contextmanager, nullcontext
from dataclasses import dataclass from dataclasses import dataclass, replace
from typing import TYPE_CHECKING, Any, Callable, Dict, Generator, List, Sequence, Union from typing import (
TYPE_CHECKING,
Any,
Callable,
Dict,
Generator,
List,
Optional,
Sequence,
Union,
)
import torch import torch
from sglang.srt.layers.dp_attention import set_dp_buffer_len from sglang.srt.layers.dp_attention import set_dp_buffer_len
from sglang.srt.model_executor.forward_context import (
forward_context,
get_forward_context,
)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import ForwardContext
_ENABLE_PROFILE = bool(int(os.environ.get("SGLANG_OPERATIONS_ENABLE_PROFILE", "0"))) _ENABLE_PROFILE = bool(int(os.environ.get("SGLANG_OPERATIONS_ENABLE_PROFILE", "0")))
@@ -39,10 +54,15 @@ def execute_overlapped_operations(
assert delta_stage_a == 0 assert delta_stage_a == 0
delta_stage = delta_stage_b delta_stage = delta_stage_b
# Each TBO child sub-batch dispatches against its own per-child backend
# (children[i] has metadata init'd for sub-batch i; the parent's primary
# has metadata for the full pre-split batch).
child_ctx_a, child_ctx_b = _resolve_tbo_child_contexts()
stages_a = _convert_operations_to_stages(operations_a) stages_a = _convert_operations_to_stages(operations_a)
stages_b = _convert_operations_to_stages(operations_b) stages_b = _convert_operations_to_stages(operations_b)
executor_a = _StageExecutor("a", stages_a, inputs=inputs_a) executor_a = _StageExecutor("a", stages_a, inputs=inputs_a, child_ctx=child_ctx_a)
executor_b = _StageExecutor("b", stages_b, inputs=inputs_b) executor_b = _StageExecutor("b", stages_b, inputs=inputs_b, child_ctx=child_ctx_b)
for _ in range(delta_stage): for _ in range(delta_stage):
executor_a.next() executor_a.next()
@@ -58,6 +78,25 @@ def execute_overlapped_operations(
return [executor_a.output, executor_b.output] return [executor_a.output, executor_b.output]
def _resolve_tbo_child_contexts():
"""Return (child_ctx_a, child_ctx_b) derived from the active TboAttnBackend,
or (None, None) if the active backend is not a TBO dispatcher (e.g. a
backend that handles TBO splitting internally like DeepSeek MHA's
_resolve_attn_backend path)."""
# Lazy import to avoid circular dependency at module load time.
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
ctx = get_forward_context()
backend = ctx.attn_backend
if not isinstance(backend, TboAttnBackend):
return None, None
child_a, child_b = backend.children
return (
replace(ctx, attn_backend=child_a),
replace(ctx, attn_backend=child_b),
)
class YieldOperation: class YieldOperation:
pass pass
@@ -73,12 +112,23 @@ Stage = List[ExecutionOperation]
class _StageExecutor: class _StageExecutor:
def __init__(self, debug_name: str, stages: List[Stage], inputs: dict): def __init__(
self,
debug_name: str,
stages: List[Stage],
inputs: dict,
child_ctx: Optional["ForwardContext"] = None,
):
self._debug_name = debug_name self._debug_name = debug_name
self._stages = stages self._stages = stages
self._index = 0 self._index = 0
self._stage_state = _StateDict() self._stage_state = _StateDict()
self._stage_output = inputs self._stage_output = inputs
# When set, every next() runs inside this ForwardContext so that
# get_attn_backend() inside RadixAttention.forward resolves to the
# per-child backend (with sub-batch metadata) instead of the TBO
# parent's primary.
self._child_ctx = child_ctx
# handling DP attention # handling DP attention
forward_batch: ForwardBatch = inputs["forward_batch"] forward_batch: ForwardBatch = inputs["forward_batch"]
@@ -102,7 +152,12 @@ class _StageExecutor:
self._global_num_tokens, self._global_num_tokens,
) )
with _annotate_region(debug_name=f"{self._debug_name}{self._index}"): ctx_mgr = (
forward_context(self._child_ctx)
if self._child_ctx is not None
else nullcontext()
)
with ctx_mgr, _annotate_region(debug_name=f"{self._debug_name}{self._index}"):
for op in stage: for op in stage:
with _annotate_region(debug_name=op.debug_name): with _annotate_region(debug_name=op.debug_name):
self._stage_output = op.fn( self._stage_output = op.fn(
@@ -14,7 +14,6 @@ from sglang.srt.batch_overlap.operations import (
) )
from sglang.srt.batch_overlap.operations_strategy import OperationsStrategy from sglang.srt.batch_overlap.operations_strategy import OperationsStrategy
from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers import deep_gemm_wrapper
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.communicator import ( from sglang.srt.layers.communicator import (
CommunicateContext, CommunicateContext,
CommunicateSummableTensorPairFn, CommunicateSummableTensorPairFn,
@@ -40,6 +39,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardMode, ForwardMode,
compute_position, compute_position,
) )
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip
@@ -508,8 +508,10 @@ class TboForwardBatchPreparer:
f"forward_mode={batch.forward_mode}" f"forward_mode={batch.forward_mode}"
) )
assert isinstance(batch.attn_backend, TboAttnBackend) # Sanity check: the global attn_backend should be a TboAttnBackend
attn_backend_child_a, attn_backend_child_b = batch.attn_backend.children # whose children handle the two halves.
attn_backend = get_attn_backend()
assert isinstance(attn_backend, TboAttnBackend)
[out_num_token_non_padded_a, out_num_token_non_padded_b] = ( [out_num_token_non_padded_a, out_num_token_non_padded_b] = (
tbo_children_num_token_non_padded tbo_children_num_token_non_padded
@@ -525,7 +527,6 @@ class TboForwardBatchPreparer:
if is_enable_two_chunk if is_enable_two_chunk
else batch.tbo_split_seq_index else batch.tbo_split_seq_index
), ),
output_attn_backend=attn_backend_child_a,
out_num_token_non_padded=out_num_token_non_padded_a, out_num_token_non_padded=out_num_token_non_padded_a,
) )
child_b = cls.filter_batch( child_b = cls.filter_batch(
@@ -534,7 +535,6 @@ class TboForwardBatchPreparer:
end_token_index=batch.input_ids.shape[0], end_token_index=batch.input_ids.shape[0],
start_seq_index=batch.tbo_split_seq_index, start_seq_index=batch.tbo_split_seq_index,
end_seq_index=batch.batch_size, end_seq_index=batch.batch_size,
output_attn_backend=attn_backend_child_b,
out_num_token_non_padded=out_num_token_non_padded_b, out_num_token_non_padded=out_num_token_non_padded_b,
) )
@@ -620,7 +620,6 @@ class TboForwardBatchPreparer:
end_token_index: int, end_token_index: int,
start_seq_index: int, start_seq_index: int,
end_seq_index: int, end_seq_index: int,
output_attn_backend: AttentionBackend,
out_num_token_non_padded: torch.Tensor, out_num_token_non_padded: torch.Tensor,
): ):
assert ( assert (
@@ -692,8 +691,6 @@ class TboForwardBatchPreparer:
"is_extend_in_batch", "is_extend_in_batch",
"all_extend_in_batch", "all_extend_in_batch",
"return_logprob", "return_logprob",
"req_to_token_pool",
"token_to_kv_pool",
"can_run_dp_cuda_graph", "can_run_dp_cuda_graph",
"dp_padding_mode", "dp_padding_mode",
"global_forward_mode", "global_forward_mode",
@@ -743,7 +740,6 @@ class TboForwardBatchPreparer:
else None else None
), ),
extend_num_tokens=extend_num_tokens, extend_num_tokens=extend_num_tokens,
attn_backend=output_attn_backend,
num_token_non_padded=out_num_token_non_padded, num_token_non_padded=out_num_token_non_padded,
# TODO: handle it when we need TBO + DeepSeek V3.2 # TODO: handle it when we need TBO + DeepSeek V3.2
num_token_non_padded_cpu=None, num_token_non_padded_cpu=None,
@@ -264,11 +264,11 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
if not self.use_mla: if not self.use_mla:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
else: else:
forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.token_to_kv_pool.set_mla_kv_buffer(
layer, layer,
cache_loc, cache_loc,
k, k,
@@ -357,9 +357,7 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
can_run_tbo=forward_batch.can_run_tbo, can_run_tbo=forward_batch.can_run_tbo,
) )
if not self.use_mla: if not self.use_mla:
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer( key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
layer.layer_id
)
key_cache = key_cache.view( key_cache = key_cache.view(
-1, self.page_size, layer.tp_k_head_num, layer.head_dim -1, self.page_size, layer.tp_k_head_num, layer.head_dim
@@ -555,9 +553,9 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
return output, lse return output, lse
return output return output
else: else:
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer( kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
layer.layer_id q.dtype
).to(q.dtype) )
k_rope = kv_cache[:, :, layer.v_head_dim :] k_rope = kv_cache[:, :, layer.v_head_dim :]
c_kv = kv_cache[:, :, : layer.v_head_dim] c_kv = kv_cache[:, :, : layer.v_head_dim]
k_rope_cache = k_rope.view( k_rope_cache = k_rope.view(
@@ -657,11 +655,11 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
if not self.use_mla: if not self.use_mla:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
else: else:
forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.token_to_kv_pool.set_mla_kv_buffer(
layer, layer,
cache_loc, cache_loc,
k, k,
@@ -710,9 +708,7 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
can_run_tbo=forward_batch.can_run_tbo, can_run_tbo=forward_batch.can_run_tbo,
) )
if not self.use_mla: if not self.use_mla:
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer( key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
layer.layer_id
)
key_cache = key_cache.view( key_cache = key_cache.view(
-1, self.page_size, layer.tp_k_head_num, layer.head_dim -1, self.page_size, layer.tp_k_head_num, layer.head_dim
) )
@@ -831,9 +827,7 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
else: else:
o = result o = result
else: else:
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id).to( kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(q.dtype)
q.dtype
)
k_rope = kv_cache[:, :, layer.v_head_dim :] k_rope = kv_cache[:, :, layer.v_head_dim :]
c_kv = kv_cache[:, :, : layer.v_head_dim] c_kv = kv_cache[:, :, : layer.v_head_dim]
k_rope_cache = k_rope.view( k_rope_cache = k_rope.view(
@@ -208,7 +208,9 @@ class AscendAttnMaskBuilder:
return attn_mask return attn_mask
def _cp_allgather_and_save_kv_npu(forward_batch, layer, k, v, cp_size): def _cp_allgather_and_save_kv_npu(
forward_batch, layer, k, v, cp_size, token_to_kv_pool
):
"""NPU-compatible CP KV all-gather with merged K/V communication. """NPU-compatible CP KV all-gather with merged K/V communication.
Merges K and V along the feature dimension so only one all-gather is Merges K and V along the feature dimension so only one all-gather is
@@ -243,7 +245,7 @@ def _cp_allgather_and_save_kv_npu(forward_batch, layer, k, v, cp_size):
key_cache_full = kv_full[..., :k_feat_size].reshape(-1, *k_tail) key_cache_full = kv_full[..., :k_feat_size].reshape(-1, *k_tail)
value_cache_full = kv_full[..., k_feat_size:].reshape(-1, *v_tail) value_cache_full = kv_full[..., k_feat_size:].reshape(-1, *v_tail)
forward_batch.token_to_kv_pool.set_kv_buffer( token_to_kv_pool.set_kv_buffer(
layer, layer,
cache_loc, cache_loc,
key_cache_full, key_cache_full,
@@ -287,6 +289,10 @@ class AscendAttnBackend(AttentionBackend):
self.native_attn = AscendTorchNativeAttnBackend() self.native_attn = AscendTorchNativeAttnBackend()
self.graph_metadata = {} self.graph_metadata = {}
self.max_context_len = model_runner.model_config.context_len self.max_context_len = model_runner.model_config.context_len
# Pool refs — captured at construction so they survive deletion of the
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.req_to_token = model_runner.req_to_token_pool.req_to_token self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.graph_mode = False self.graph_mode = False
self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False") self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False")
@@ -357,7 +363,7 @@ class AscendAttnBackend(AttentionBackend):
): ):
seq_lens_max += self.speculative_step_id + 1 seq_lens_max += self.speculative_step_id + 1
self.forward_metadata.block_tables = ( self.forward_metadata.block_tables = (
forward_batch.req_to_token_pool.req_to_token[ self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, :seq_lens_max forward_batch.req_pool_indices, :seq_lens_max
][:, :: self.page_size] ][:, :: self.page_size]
// self.page_size // self.page_size
@@ -366,7 +372,7 @@ class AscendAttnBackend(AttentionBackend):
self.forward_metadata.block_tables_swa = ( self.forward_metadata.block_tables_swa = (
( (
self.full_to_swa_index_mapping[ self.full_to_swa_index_mapping[
forward_batch.req_to_token_pool.req_to_token[ self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, :seq_lens_max forward_batch.req_pool_indices, :seq_lens_max
] ]
][:, :: self.page_size] ][:, :: self.page_size]
@@ -421,7 +427,7 @@ class AscendAttnBackend(AttentionBackend):
for req_idx, seq_len in zip( for req_idx, seq_len in zip(
forward_batch.req_pool_indices.tolist(), seq_prefix_lens forward_batch.req_pool_indices.tolist(), seq_prefix_lens
): ):
req_indices = forward_batch.req_to_token_pool.req_to_token[req_idx] req_indices = self.req_to_token_pool.req_to_token[req_idx]
req_prefix_block_tables = ( req_prefix_block_tables = (
req_indices[:seq_len][:: self.page_size] // self.page_size req_indices[:seq_len][:: self.page_size] // self.page_size
) )
@@ -883,11 +889,11 @@ class AscendAttnBackend(AttentionBackend):
if save_kv_cache: if save_kv_cache:
k = k.view(-1, layer.tp_k_head_num, self.kv_lora_rank) k = k.view(-1, layer.tp_k_head_num, self.kv_lora_rank)
k_rope = k_rope.view(-1, layer.tp_k_head_num, self.qk_rope_head_dim) k_rope = k_rope.view(-1, layer.tp_k_head_num, self.qk_rope_head_dim)
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, k_rope layer, forward_batch.out_cache_loc, k, k_rope
) )
q_nope, q_pe = q, q_rope q_nope, q_pe = q, q_rope
k_nope, k_pe = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id) k_nope, k_pe = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
if is_prefill: if is_prefill:
if self.forward_metadata.actual_seq_lengths_q is not None: if self.forward_metadata.actual_seq_lengths_q is not None:
@@ -1041,7 +1047,12 @@ class AscendAttnBackend(AttentionBackend):
if is_cp_mode: if is_cp_mode:
# All-gather K/V from all CP ranks and write full sequence to KV pool # All-gather K/V from all CP ranks and write full sequence to KV pool
_cp_allgather_and_save_kv_npu( _cp_allgather_and_save_kv_npu(
forward_batch, layer, k, v, self.attn_cp_size forward_batch,
layer,
k,
v,
self.attn_cp_size,
self.token_to_kv_pool,
) )
else: else:
# support cross attention # support cross attention
@@ -1050,10 +1061,10 @@ class AscendAttnBackend(AttentionBackend):
if not layer.is_cross_attention if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
forward_batch.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id) v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
if sinks is not None: if sinks is not None:
# Use SWA block tables if hybrid SWA is enabled for this layer # Use SWA block tables if hybrid SWA is enabled for this layer
@@ -1200,7 +1211,7 @@ class AscendAttnBackend(AttentionBackend):
o_, o_,
k_cache.view(-1, layer.tp_k_head_num, layer.qk_head_dim), k_cache.view(-1, layer.tp_k_head_num, layer.qk_head_dim),
v_cache.view(-1, layer.tp_v_head_num, layer.v_head_dim), v_cache.view(-1, layer.tp_v_head_num, layer.v_head_dim),
forward_batch.req_to_token_pool.req_to_token, self.req_to_token_pool.req_to_token,
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.seq_lens, forward_batch.seq_lens,
forward_batch.extend_prefix_lens, forward_batch.extend_prefix_lens,
@@ -1223,10 +1234,8 @@ class AscendAttnBackend(AttentionBackend):
if layer.qk_head_dim == layer.v_head_dim: if layer.qk_head_dim == layer.v_head_dim:
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim) q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
k_buffer = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_buffer = forward_batch.token_to_kv_pool.get_value_buffer( v_buffer = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
layer.layer_id
)
kv_cached = torch.index_select( kv_cached = torch.index_select(
k_buffer, 0, self.forward_metadata.flatten_prefix_block_tables k_buffer, 0, self.forward_metadata.flatten_prefix_block_tables
) )
@@ -1335,10 +1344,8 @@ class AscendAttnBackend(AttentionBackend):
) )
# 2nd, load history kvcache(kv_a and k_pe) and calculate k_nope # 2nd, load history kvcache(kv_a and k_pe) and calculate k_nope
k_buffer = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_buffer = forward_batch.token_to_kv_pool.get_value_buffer( v_buffer = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
layer.layer_id
)
kv_cached = torch.index_select( kv_cached = torch.index_select(
k_buffer, 0, self.forward_metadata.flatten_prefix_block_tables k_buffer, 0, self.forward_metadata.flatten_prefix_block_tables
) )
@@ -1427,7 +1434,7 @@ class AscendAttnBackend(AttentionBackend):
kv_lora_rank = k.shape[-1] - self.qk_rope_head_dim kv_lora_rank = k.shape[-1] - self.qk_rope_head_dim
kv_c, k_rope = k.split([kv_lora_rank, self.qk_rope_head_dim], dim=-1) kv_c, k_rope = k.split([kv_lora_rank, self.qk_rope_head_dim], dim=-1)
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, kv_c, k_rope layer, forward_batch.out_cache_loc, kv_c, k_rope
) )
attn_output = q.new_empty( attn_output = q.new_empty(
@@ -1435,17 +1442,15 @@ class AscendAttnBackend(AttentionBackend):
) )
use_gqa = layer.tp_q_head_num != layer.tp_k_head_num use_gqa = layer.tp_q_head_num != layer.tp_k_head_num
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_cache = forward_batch.token_to_kv_pool.get_value_buffer( v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
layer.layer_id
)
kv_cache = torch.cat([k_cache, v_cache], dim=-1) kv_cache = torch.cat([k_cache, v_cache], dim=-1)
attn_output = self.native_attn.run_sdpa_forward_extend( attn_output = self.native_attn.run_sdpa_forward_extend(
q, q,
attn_output, attn_output,
kv_cache.view(-1, layer.tp_k_head_num, layer.qk_head_dim), kv_cache.view(-1, layer.tp_k_head_num, layer.qk_head_dim),
k_cache.view(-1, layer.tp_v_head_num, layer.v_head_dim), k_cache.view(-1, layer.tp_v_head_num, layer.v_head_dim),
forward_batch.req_to_token_pool.req_to_token, self.req_to_token_pool.req_to_token,
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.seq_lens, forward_batch.seq_lens,
forward_batch.extend_prefix_lens, forward_batch.extend_prefix_lens,
@@ -1514,12 +1519,12 @@ class AscendAttnBackend(AttentionBackend):
topk_indices: Optional[torch.Tensor] = None, topk_indices: Optional[torch.Tensor] = None,
): ):
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v layer, forward_batch.out_cache_loc, k, v
) )
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id) v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
query = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim) query = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
if self.forward_metadata.seq_lens_cpu_int is None: if self.forward_metadata.seq_lens_cpu_int is None:
@@ -1574,21 +1579,21 @@ class AscendAttnBackend(AttentionBackend):
if self.use_mla: if self.use_mla:
k = k.view(-1, layer.tp_k_head_num, self.kv_lora_rank) k = k.view(-1, layer.tp_k_head_num, self.kv_lora_rank)
k_rope = k_rope.view(-1, layer.tp_k_head_num, self.qk_rope_head_dim) k_rope = k_rope.view(-1, layer.tp_k_head_num, self.qk_rope_head_dim)
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, k_rope layer, forward_batch.out_cache_loc, k, k_rope
) )
else: else:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v layer, forward_batch.out_cache_loc, k, v
) )
if not self.use_mla: if not self.use_mla:
k_cache = forward_batch.token_to_kv_pool.get_key_buffer( k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).view(
layer.layer_id -1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
).view(-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim) )
v_cache = forward_batch.token_to_kv_pool.get_value_buffer( v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id).view(
layer.layer_id -1, self.page_size, layer.tp_v_head_num * layer.v_head_dim
).view(-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim) )
query = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim).contiguous() query = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim).contiguous()
if not self.graph_mode: if not self.graph_mode:
num_token_padding = query.shape[0] num_token_padding = query.shape[0]
@@ -1642,7 +1647,7 @@ class AscendAttnBackend(AttentionBackend):
) )
return attn_output return attn_output
else: else:
c_kv, k_rope = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id) c_kv, k_rope = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
if is_fia_nz(): if is_fia_nz():
k_rope_cache = _reshape_kv_for_fia_nz( k_rope_cache = _reshape_kv_for_fia_nz(
k_rope, layer.tp_k_head_num, self.qk_rope_head_dim, self.page_size k_rope, layer.tp_k_head_num, self.qk_rope_head_dim, self.page_size
@@ -1756,17 +1761,17 @@ class AscendAttnBackend(AttentionBackend):
if self.use_mla: if self.use_mla:
k = k.view(-1, layer.tp_k_head_num, self.kv_lora_rank) k = k.view(-1, layer.tp_k_head_num, self.kv_lora_rank)
k_rope = k_rope.view(-1, layer.tp_k_head_num, self.qk_rope_head_dim) k_rope = k_rope.view(-1, layer.tp_k_head_num, self.qk_rope_head_dim)
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, k_rope layer, forward_batch.out_cache_loc, k, k_rope
) )
else: else:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v layer, forward_batch.out_cache_loc, k, v
) )
if sinks is not None: if sinks is not None:
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id) v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
# Use SWA block tables if hybrid SWA is enabled for this layer # Use SWA block tables if hybrid SWA is enabled for this layer
if self.is_hybrid_swa and layer.sliding_window_size != -1: if self.is_hybrid_swa and layer.sliding_window_size != -1:
@@ -1788,12 +1793,12 @@ class AscendAttnBackend(AttentionBackend):
return attn_out return attn_out
if not self.use_mla: if not self.use_mla:
k_cache = forward_batch.token_to_kv_pool.get_key_buffer( k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).view(
layer.layer_id -1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
).view(-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim) )
v_cache = forward_batch.token_to_kv_pool.get_value_buffer( v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id).view(
layer.layer_id -1, self.page_size, layer.tp_v_head_num * layer.v_head_dim
).view(-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim) )
query = q.reshape(-1, 1, layer.tp_q_head_num * layer.qk_head_dim) query = q.reshape(-1, 1, layer.tp_q_head_num * layer.qk_head_dim)
if self.forward_metadata.seq_lens_cpu_int is None: if self.forward_metadata.seq_lens_cpu_int is None:
actual_seq_len_kv = self.forward_metadata.seq_lens_cpu_list actual_seq_len_kv = self.forward_metadata.seq_lens_cpu_list
@@ -1836,7 +1841,7 @@ class AscendAttnBackend(AttentionBackend):
) )
return output.view(num_tokens, layer.tp_q_head_num * layer.v_head_dim) return output.view(num_tokens, layer.tp_q_head_num * layer.v_head_dim)
else: else:
c_kv, k_rope = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id) c_kv, k_rope = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
if is_fia_nz(): if is_fia_nz():
k_rope_cache = _reshape_kv_for_fia_nz( k_rope_cache = _reshape_kv_for_fia_nz(
k_rope, layer.tp_k_head_num, self.qk_rope_head_dim, self.page_size k_rope, layer.tp_k_head_num, self.qk_rope_head_dim, self.page_size
@@ -1976,10 +1981,10 @@ class AscendAttnBackend(AttentionBackend):
if not layer.is_cross_attention if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
forward_batch.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
num_tokens = q.shape[0] num_tokens = q.shape[0]
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id) v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
if sinks is not None: if sinks is not None:
# Use SWA block tables if hybrid SWA is enabled for this layer # Use SWA block tables if hybrid SWA is enabled for this layer
@@ -2098,7 +2103,7 @@ class AscendAttnBackend(AttentionBackend):
o_, o_,
k_cache.view(-1, layer.tp_k_head_num, layer.qk_head_dim), k_cache.view(-1, layer.tp_k_head_num, layer.qk_head_dim),
v_cache.view(-1, layer.tp_v_head_num, layer.v_head_dim), v_cache.view(-1, layer.tp_v_head_num, layer.v_head_dim),
forward_batch.req_to_token_pool.req_to_token, self.req_to_token_pool.req_to_token,
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.seq_lens, forward_batch.seq_lens,
forward_batch.encoder_lens, forward_batch.encoder_lens,
@@ -2112,12 +2117,12 @@ class AscendAttnBackend(AttentionBackend):
return attn_output.view(num_tokens, layer.tp_q_head_num * layer.v_head_dim) return attn_output.view(num_tokens, layer.tp_q_head_num * layer.v_head_dim)
else: else:
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, k_rope layer, forward_batch.out_cache_loc, k, k_rope
) )
num_tokens = q.shape[0] num_tokens = q.shape[0]
kv_c = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) kv_c = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
k_pe = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id) k_pe = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
if self.use_fia and (layer.tp_q_head_num // layer.tp_k_head_num) >= 8: if self.use_fia and (layer.tp_q_head_num // layer.tp_k_head_num) >= 8:
"""layer.tp_q_head_num // layer.tp_k_head_num < 8 will support in the later version of CANN""" """layer.tp_q_head_num // layer.tp_k_head_num < 8 will support in the later version of CANN"""
@@ -2218,11 +2223,11 @@ class AscendAttnBackend(AttentionBackend):
"3. When the environment variable ASCEND_USE_FIA is set to 0 and qk_head_dim exceeds 128 on Ascend NPU devices." "3. When the environment variable ASCEND_USE_FIA is set to 0 and qk_head_dim exceeds 128 on Ascend NPU devices."
) )
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v layer, forward_batch.out_cache_loc, k, v
) )
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id) v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
num_block, block_size, _, _ = k_cache.shape num_block, block_size, _, _ = k_cache.shape
key = k_cache.view(num_block, block_size, -1) key = k_cache.view(num_block, block_size, -1)
value = v_cache.view(num_block, block_size, -1) value = v_cache.view(num_block, block_size, -1)
@@ -6,6 +6,10 @@ import torch
import torch.nn.functional as F import torch.nn.functional as F
from sglang.srt.hardware_backend.npu.utils import npu_format_cast from sglang.srt.hardware_backend.npu.utils import npu_format_cast
from sglang.srt.model_executor.forward_context import (
get_attn_backend,
get_token_to_kv_pool,
)
from sglang.srt.utils import get_bool_env_var from sglang.srt.utils import get_bool_env_var
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -253,7 +257,7 @@ class NPUFusedMLAPreprocess(torch.nn.Module):
return cos, sin return cos, sin
def get_kv_cache_and_cache_idx(self, forward_batch): def get_kv_cache_and_cache_idx(self, forward_batch):
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer(self.layer_id) k_cache, v_cache = get_token_to_kv_pool().get_kv_buffer(self.layer_id)
slot_mapping = forward_batch.out_cache_loc.to(dtype=torch.int32) slot_mapping = forward_batch.out_cache_loc.to(dtype=torch.int32)
return k_cache, v_cache, slot_mapping return k_cache, v_cache, slot_mapping
@@ -314,15 +318,15 @@ class NPUFusedMLAPreprocess(torch.nn.Module):
cache_mode = "PA_NZ" if is_fia_nz() else "PA_BNSD" cache_mode = "PA_NZ" if is_fia_nz() else "PA_BNSD"
self.kvCache = self.kvCache.view( self.kvCache = self.kvCache.view(
-1, -1,
forward_batch.attn_backend.page_size, get_attn_backend().page_size,
1, 1,
forward_batch.attn_backend.kv_lora_rank, get_attn_backend().kv_lora_rank,
) )
self.kvCacheRope = self.kvCacheRope.view( self.kvCacheRope = self.kvCacheRope.view(
-1, -1,
forward_batch.attn_backend.page_size, get_attn_backend().page_size,
1, 1,
forward_batch.attn_backend.qk_rope_head_dim, get_attn_backend().qk_rope_head_dim,
) )
k_rope, k_nope, _, _ = torch.ops.npu.npu_kv_rmsnorm_rope_cache( k_rope, k_nope, _, _ = torch.ops.npu.npu_kv_rmsnorm_rope_cache(
latent_cache, latent_cache,
@@ -16,6 +16,7 @@ from sglang.srt.layers.attention.dsa.utils import (
dsa_use_prefill_cp, dsa_use_prefill_cp,
) )
from sglang.srt.layers.communicator import ScatterMode, get_attn_tp_context from sglang.srt.layers.communicator import ScatterMode, get_attn_tp_context
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@@ -88,9 +89,7 @@ def forward_mha_prepare_npu(
) )
q_pe = q_pe.reshape(B, -1, m.qk_rope_head_dim) q_pe = q_pe.reshape(B, -1, m.qk_rope_head_dim)
ckv_cache, k_rope_cache = forward_batch.token_to_kv_pool.get_kv_buffer( ckv_cache, k_rope_cache = get_token_to_kv_pool().get_kv_buffer(m.layer_id)
m.layer_id
)
_, _, k_pe, kv_a = torch_npu.npu_kv_rmsnorm_rope_cache( _, _, k_pe, kv_a = torch_npu.npu_kv_rmsnorm_rope_cache(
latent_cache.view(-1, 1, 1, m.kv_lora_rank + m.qk_rope_head_dim), # bnsd latent_cache.view(-1, 1, 1, m.kv_lora_rank + m.qk_rope_head_dim), # bnsd
m.kv_a_layernorm.weight, m.kv_a_layernorm.weight,
@@ -115,7 +114,7 @@ def forward_mha_prepare_npu(
if m.rotary_emb is not None: if m.rotary_emb is not None:
q_pe, k_pe = m.rotary_emb(positions, q_pe, k_pe) q_pe, k_pe = m.rotary_emb(positions, q_pe, k_pe)
# this is for model kimi-vl-a3B-instruct # this is for model kimi-vl-a3B-instruct
forward_batch.token_to_kv_pool.set_kv_buffer( get_token_to_kv_pool().set_kv_buffer(
m, forward_batch.out_cache_loc, kv_a.unsqueeze(1), k_pe m, forward_batch.out_cache_loc, kv_a.unsqueeze(1), k_pe
) )
@@ -204,6 +204,11 @@ class AiterAttnBackend(AttentionBackend):
model_runner, self model_runner, self
) )
# Pool refs — captured at construction so they survive deletion of the
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
# sliding window attention # sliding window attention
self.use_sliding_window_kv_pool = ( self.use_sliding_window_kv_pool = (
isinstance(model_runner.token_to_kv_pool, SWAKVPool) isinstance(model_runner.token_to_kv_pool, SWAKVPool)
@@ -211,7 +216,6 @@ class AiterAttnBackend(AttentionBackend):
) )
if self.use_sliding_window_kv_pool: if self.use_sliding_window_kv_pool:
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.use_triton_unified_attention = True self.use_triton_unified_attention = True
else: else:
self.use_triton_unified_attention = get_bool_env_var( self.use_triton_unified_attention = get_bool_env_var(
@@ -2355,8 +2359,8 @@ class AiterAttnBackend(AttentionBackend):
self.use_triton_unified_attention self.use_triton_unified_attention
and self.use_sliding_window_kv_pool and self.use_sliding_window_kv_pool
): ):
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = self.token_to_kv_pool
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer( k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer(
layer.layer_id layer.layer_id
) )
slot_mapping_swa = token_to_kv_pool.full_to_swa_index_mapping slot_mapping_swa = token_to_kv_pool.full_to_swa_index_mapping
@@ -2380,9 +2384,9 @@ class AiterAttnBackend(AttentionBackend):
v_scale=v_descale, v_scale=v_descale,
) )
elif self.use_mla: elif self.use_mla:
forward_batch.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
else: else:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, k_descale, v_descale layer, cache_loc, k, v, k_descale, v_descale
) )
@@ -2392,8 +2396,8 @@ class AiterAttnBackend(AttentionBackend):
kv_indptr = self.forward_metadata.kv_indptr kv_indptr = self.forward_metadata.kv_indptr
kv_indices = self.forward_metadata.kv_indices kv_indices = self.forward_metadata.kv_indices
qo_indptr = self.forward_metadata.qo_indptr qo_indptr = self.forward_metadata.qo_indptr
K_Buffer = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) K_Buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
V_Buffer = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id) V_Buffer = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
kv_lora_rank = V_Buffer.shape[-1] kv_lora_rank = V_Buffer.shape[-1]
qk_rope_head_dim = K_Buffer.shape[-1] - kv_lora_rank qk_rope_head_dim = K_Buffer.shape[-1] - kv_lora_rank
qk_nope_head_dim = k.shape[-1] - qk_rope_head_dim qk_nope_head_dim = k.shape[-1] - qk_rope_head_dim
@@ -2646,7 +2650,7 @@ class AiterAttnBackend(AttentionBackend):
self._use_unified_verify self._use_unified_verify
and forward_batch.forward_mode.is_target_verify() and forward_batch.forward_mode.is_target_verify()
): ):
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer( k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer(
layer.layer_id layer.layer_id
) )
page_table = self.forward_metadata.kv_indices page_table = self.forward_metadata.kv_indices
@@ -2705,8 +2709,8 @@ class AiterAttnBackend(AttentionBackend):
k.contiguous(), k.contiguous(),
v.contiguous(), v.contiguous(),
o.view(-1, layer.tp_q_head_num, layer.v_head_dim), o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), self.token_to_kv_pool.get_key_buffer(layer.layer_id),
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id), self.token_to_kv_pool.get_value_buffer(layer.layer_id),
self.forward_metadata.qo_indptr, self.forward_metadata.qo_indptr,
self.forward_metadata.kv_indptr, self.forward_metadata.kv_indptr,
self.forward_metadata.kv_indices, self.forward_metadata.kv_indices,
@@ -2721,9 +2725,7 @@ class AiterAttnBackend(AttentionBackend):
) )
return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) return o.view(-1, layer.tp_q_head_num * layer.v_head_dim)
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer( k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
layer.layer_id
)
bs0 = forward_batch.batch_size + 1 bs0 = forward_batch.batch_size + 1
@@ -2798,10 +2800,8 @@ class AiterAttnBackend(AttentionBackend):
# use standard set_kv_buffer, as they lack SWA-specific attributes # use standard set_kv_buffer, as they lack SWA-specific attributes
# like full_to_swa_index_mapping. # like full_to_swa_index_mapping.
if self.use_triton_unified_attention and self.use_sliding_window_kv_pool: if self.use_triton_unified_attention and self.use_sliding_window_kv_pool:
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = self.token_to_kv_pool
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer( k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
layer.layer_id
)
slot_mapping_swa = token_to_kv_pool.full_to_swa_index_mapping slot_mapping_swa = token_to_kv_pool.full_to_swa_index_mapping
launch_reshape_and_cache_flash( launch_reshape_and_cache_flash(
@@ -2822,7 +2822,7 @@ class AiterAttnBackend(AttentionBackend):
# [PATCH] FP8 non-SWA: use launch_reshape_and_cache_flash to # [PATCH] FP8 non-SWA: use launch_reshape_and_cache_flash to
# fuse bf16→fp8 cast + paged write in one Triton kernel, # fuse bf16→fp8 cast + paged write in one Triton kernel,
# eliminating separate float8_copy + store_kvcache overhead. # eliminating separate float8_copy + store_kvcache overhead.
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = self.token_to_kv_pool
k_cache, v_cache = token_to_kv_pool.get_kv_buffer(layer.layer_id) k_cache, v_cache = token_to_kv_pool.get_kv_buffer(layer.layer_id)
launch_reshape_and_cache_flash( launch_reshape_and_cache_flash(
k.view(-1, layer.tp_k_head_num, layer.qk_head_dim), k.view(-1, layer.tp_k_head_num, layer.qk_head_dim),
@@ -2836,12 +2836,12 @@ class AiterAttnBackend(AttentionBackend):
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
) )
else: else:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v layer, forward_batch.out_cache_loc, k, v
) )
if self.use_mla: if self.use_mla:
k_buffer = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
work_metadata = self.forward_metadata.work_metadata work_metadata = self.forward_metadata.work_metadata
work_indptr = self.forward_metadata.work_indptr work_indptr = self.forward_metadata.work_indptr
@@ -2878,9 +2878,7 @@ class AiterAttnBackend(AttentionBackend):
else: else:
self.logits_soft_cap = layer.logit_cap self.logits_soft_cap = layer.logit_cap
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer( k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
layer.layer_id
)
if layer.qk_head_dim != layer.v_head_dim: if layer.qk_head_dim != layer.v_head_dim:
o = q.new_empty( o = q.new_empty(
@@ -3185,6 +3183,7 @@ class AiterMultiStepDraftBackend:
) )
self.device = model_runner.device self.device = model_runner.device
# Cached variables for generate_draft_decode_kv_indices # Cached variables for generate_draft_decode_kv_indices
self.req_to_token_pool = model_runner.req_to_token_pool
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1] self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
self.page_size = model_runner.server_args.page_size self.page_size = model_runner.server_args.page_size
@@ -3199,7 +3198,7 @@ class AiterMultiStepDraftBackend:
(self.speculative_num_steps, num_seqs, self.topk) (self.speculative_num_steps, num_seqs, self.topk)
]( ](
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.req_to_token_pool.req_to_token, self.req_to_token_pool.req_to_token,
forward_batch.seq_lens, forward_batch.seq_lens,
kv_indices_buffer, kv_indices_buffer,
self.kv_indptr, self.kv_indptr,
@@ -241,14 +241,14 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
assert v is not None assert v is not None
if save_kv_cache: if save_kv_cache:
if k_rope is not None: if k_rope is not None:
forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.token_to_kv_pool.set_mla_kv_buffer(
layer, layer,
cache_loc, cache_loc,
k, k,
k_rope, k_rope,
) )
else: else:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, layer,
cache_loc, cache_loc,
k, k,
@@ -269,7 +269,7 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
q_nope = q_nope.to(self.q_data_type) q_nope = q_nope.to(self.q_data_type)
q_rope = q_rope.to(self.q_data_type) q_rope = q_rope.to(self.q_data_type)
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
o = cutlass_mla_decode( o = cutlass_mla_decode(
q_nope=q_nope, q_nope=q_nope,
@@ -354,8 +354,15 @@ class DeepseekV4AttnBackend(
self.page_size = model_runner.page_size self.page_size = model_runner.page_size
assert self.page_size == 256, "the system hardcodes page_size=256" assert self.page_size == 256, "the system hardcodes page_size=256"
self.req_to_token = model_runner.req_to_token_pool.req_to_token # Pool refs — captured at construction so they survive deletion of the
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool: DeepSeekV4TokenToKVPool = model_runner.token_to_kv_pool self.token_to_kv_pool: DeepSeekV4TokenToKVPool = model_runner.token_to_kv_pool
# Keep a runner ref to read live state set after backend construction
# (e.g. hisparse_coordinator is built in model_runner *after*
# init_attention_backend()).
self.model_runner = model_runner
self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.MAX_SEQ_LEN_FOR_CAPTURE = self.req_to_token.shape[1] self.MAX_SEQ_LEN_FOR_CAPTURE = self.req_to_token.shape[1]
assert isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool)
@@ -378,6 +385,12 @@ class DeepseekV4AttnBackend(
] = None ] = None
self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band
@property
def hisparse_coordinator(self):
# Live read: model_runner builds the coordinator *after*
# init_attention_backend(), so we cannot capture at __init__ time.
return self.model_runner.hisparse_coordinator
def _move_to_device(self, x: List[int]) -> torch.Tensor: def _move_to_device(self, x: List[int]) -> torch.Tensor:
pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True) pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True)
return pin_tensor.to(self.device, non_blocking=True) return pin_tensor.to(self.device, non_blocking=True)
@@ -667,7 +680,7 @@ class DeepseekV4AttnBackend(
req_pool_indices = forward_batch.req_pool_indices req_pool_indices = forward_batch.req_pool_indices
seq_lens = forward_batch.seq_lens.to(torch.int32) seq_lens = forward_batch.seq_lens.to(torch.int32)
seq_lens_cpu = forward_batch.seq_lens_cpu seq_lens_cpu = forward_batch.seq_lens_cpu
assert forward_batch.req_to_token_pool.req_to_token is self.req_to_token assert self.req_to_token_pool.req_to_token is self.req_to_token
assert self.swa_page_size % SWA_WINDOW == 0 and self.page_size % 128 == 0 assert self.swa_page_size % SWA_WINDOW == 0 and self.page_size % 128 == 0
assert seq_lens_cpu is not None assert seq_lens_cpu is not None
@@ -960,7 +973,7 @@ class DeepseekV4AttnBackend(
layer_id = layer.layer_id layer_id = layer.layer_id
metadata = self.forward_metadata metadata = self.forward_metadata
core_attn_metadata = metadata.core_attn_metadata core_attn_metadata = metadata.core_attn_metadata
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = self.token_to_kv_pool
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
if isinstance(core_attn_metadata, DSV4AttnMetadata): if isinstance(core_attn_metadata, DSV4AttnMetadata):
@@ -348,8 +348,15 @@ class DeepseekV4HipRadixBackend(
self.page_size = model_runner.page_size self.page_size = model_runner.page_size
assert self.page_size == 256, "the system hardcodes page_size=256" assert self.page_size == 256, "the system hardcodes page_size=256"
self.req_to_token = model_runner.req_to_token_pool.req_to_token # Pool refs — captured at construction so they survive deletion of the
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool: DeepSeekV4TokenToKVPool = model_runner.token_to_kv_pool self.token_to_kv_pool: DeepSeekV4TokenToKVPool = model_runner.token_to_kv_pool
# Keep a runner ref to read live state set after backend construction
# (e.g. hisparse_coordinator is built in model_runner *after*
# init_attention_backend()).
self.model_runner = model_runner
self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.MAX_SEQ_LEN_FOR_CAPTURE = self.req_to_token.shape[1] self.MAX_SEQ_LEN_FOR_CAPTURE = self.req_to_token.shape[1]
assert isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool)
@@ -372,6 +379,12 @@ class DeepseekV4HipRadixBackend(
] = None ] = None
self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band
@property
def hisparse_coordinator(self):
# Live read: model_runner builds the coordinator *after*
# init_attention_backend(), so we cannot capture at __init__ time.
return self.model_runner.hisparse_coordinator
def _move_to_device(self, x: List[int]) -> torch.Tensor: def _move_to_device(self, x: List[int]) -> torch.Tensor:
pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True) pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True)
return pin_tensor.to(self.device, non_blocking=True) return pin_tensor.to(self.device, non_blocking=True)
@@ -661,7 +674,7 @@ class DeepseekV4HipRadixBackend(
req_pool_indices = forward_batch.req_pool_indices req_pool_indices = forward_batch.req_pool_indices
seq_lens = forward_batch.seq_lens.to(torch.int32) seq_lens = forward_batch.seq_lens.to(torch.int32)
seq_lens_cpu = forward_batch.seq_lens_cpu seq_lens_cpu = forward_batch.seq_lens_cpu
assert forward_batch.req_to_token_pool.req_to_token is self.req_to_token assert self.req_to_token_pool.req_to_token is self.req_to_token
assert self.swa_page_size % SWA_WINDOW == 0 and self.page_size % 128 == 0 assert self.swa_page_size % SWA_WINDOW == 0 and self.page_size % 128 == 0
assert seq_lens_cpu is not None assert seq_lens_cpu is not None
@@ -954,7 +967,7 @@ class DeepseekV4HipRadixBackend(
layer_id = layer.layer_id layer_id = layer.layer_id
metadata = self.forward_metadata metadata = self.forward_metadata
core_attn_metadata = metadata.core_attn_metadata core_attn_metadata = metadata.core_attn_metadata
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = self.token_to_kv_pool
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
if isinstance(core_attn_metadata, DSV4AttnMetadata): if isinstance(core_attn_metadata, DSV4AttnMetadata):
@@ -80,6 +80,11 @@ from sglang.srt.layers.rotary_embedding import get_rope_wrapper
from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import (
get_attn_backend,
get_req_to_token_pool,
get_token_to_kv_pool,
)
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
_use_ag_after_qlora = envs.SGLANG_USE_AG_AFTER_QLORA.get() _use_ag_after_qlora = envs.SGLANG_USE_AG_AFTER_QLORA.get()
@@ -449,9 +454,9 @@ class Indexer(MultiPlatformOp):
metadata: BaseIndexerMetadata, metadata: BaseIndexerMetadata,
) -> torch.Tensor: ) -> torch.Tensor:
if TYPE_CHECKING: if TYPE_CHECKING:
assert isinstance(forward_batch.token_to_kv_pool, DSATokenToKVPool) assert isinstance(get_token_to_kv_pool(), DSATokenToKVPool)
page_size = forward_batch.token_to_kv_pool.page_size page_size = get_token_to_kv_pool().page_size
# NOTE(dark): blocksize = 64 is hardcoded in deep_gemm # NOTE(dark): blocksize = 64 is hardcoded in deep_gemm
if _is_hip: if _is_hip:
if _use_aiter_preshuffle: if _use_aiter_preshuffle:
@@ -471,7 +476,7 @@ class Indexer(MultiPlatformOp):
block_tables = metadata.get_page_table_64() block_tables = metadata.get_page_table_64()
max_seq_len = block_tables.shape[1] * page_size max_seq_len = block_tables.shape[1] * page_size
kv_cache_fp8 = forward_batch.token_to_kv_pool.get_index_k_with_scale_buffer( kv_cache_fp8 = get_token_to_kv_pool().get_index_k_with_scale_buffer(
layer_id=layer_id layer_id=layer_id
) )
@@ -624,11 +629,11 @@ class Indexer(MultiPlatformOp):
metadata: BaseIndexerMetadata, metadata: BaseIndexerMetadata,
) -> torch.Tensor: ) -> torch.Tensor:
if TYPE_CHECKING: if TYPE_CHECKING:
assert isinstance(forward_batch.token_to_kv_pool, DSATokenToKVPool) assert isinstance(get_token_to_kv_pool(), DSATokenToKVPool)
assert forward_batch.forward_mode.is_extend_without_speculative() assert forward_batch.forward_mode.is_extend_without_speculative()
page_size = forward_batch.token_to_kv_pool.page_size page_size = get_token_to_kv_pool().page_size
if _is_hip: if _is_hip:
if _use_aiter_preshuffle: if _use_aiter_preshuffle:
assert ( assert (
@@ -675,7 +680,7 @@ class Indexer(MultiPlatformOp):
indexer_seq_lens_cpu = metadata.get_indexer_seq_len_cpu() indexer_seq_lens_cpu = metadata.get_indexer_seq_len_cpu()
seq_len_sum = torch.sum(indexer_seq_lens_cpu).item() seq_len_sum = torch.sum(indexer_seq_lens_cpu).item()
max_seq_len = torch.max(indexer_seq_lens_cpu).item() max_seq_len = torch.max(indexer_seq_lens_cpu).item()
k_fp8, k_scale = forward_batch.token_to_kv_pool.get_index_k_scale_buffer( k_fp8, k_scale = get_token_to_kv_pool().get_index_k_scale_buffer(
layer_id, layer_id,
metadata.get_indexer_seq_len(), metadata.get_indexer_seq_len(),
block_tables, block_tables,
@@ -851,9 +856,9 @@ class Indexer(MultiPlatformOp):
cp_index: List[Tuple[int, int, int]] = None, cp_index: List[Tuple[int, int, int]] = None,
) -> torch.Tensor: ) -> torch.Tensor:
if TYPE_CHECKING: if TYPE_CHECKING:
assert isinstance(forward_batch.token_to_kv_pool, DSATokenToKVPool) assert isinstance(get_token_to_kv_pool(), DSATokenToKVPool)
page_size = forward_batch.token_to_kv_pool.page_size page_size = get_token_to_kv_pool().page_size
assert page_size == 64, "only support page size 64" assert page_size == 64, "only support page size 64"
assert len(weights.shape) == 3 assert len(weights.shape) == 3
weights = weights.squeeze(-1) weights = weights.squeeze(-1)
@@ -882,12 +887,12 @@ class Indexer(MultiPlatformOp):
end_seq_position += pre_chunk_offset end_seq_position += pre_chunk_offset
if offset == 0 and batch_idx != 0: if offset == 0 and batch_idx != 0:
offset += forward_batch.extend_seq_lens_cpu[batch_idx - 1] offset += forward_batch.extend_seq_lens_cpu[batch_idx - 1]
k_fp8 = forward_batch.token_to_kv_pool.get_index_k_continuous( k_fp8 = get_token_to_kv_pool().get_index_k_continuous(
layer_id, layer_id,
end_seq_position, end_seq_position,
block_tables[batch_idx], block_tables[batch_idx],
) )
k_scale = forward_batch.token_to_kv_pool.get_index_k_scale_continuous( k_scale = get_token_to_kv_pool().get_index_k_scale_continuous(
layer_id, layer_id,
end_seq_position, end_seq_position,
block_tables[batch_idx], block_tables[batch_idx],
@@ -943,12 +948,12 @@ class Indexer(MultiPlatformOp):
- forward_batch.extend_seq_lens_cpu[0] - forward_batch.extend_seq_lens_cpu[0]
+ kv_len + kv_len
) )
k_fp8 = forward_batch.token_to_kv_pool.get_index_k_continuous( k_fp8 = get_token_to_kv_pool().get_index_k_continuous(
layer_id, layer_id,
kv_len, kv_len,
block_tables[0], block_tables[0],
) )
k_scale = forward_batch.token_to_kv_pool.get_index_k_scale_continuous( k_scale = get_token_to_kv_pool().get_index_k_scale_continuous(
layer_id, layer_id,
kv_len, kv_len,
block_tables[0], block_tables[0],
@@ -999,7 +1004,7 @@ class Indexer(MultiPlatformOp):
if not _is_npu: if not _is_npu:
from sglang.srt.layers.attention.dsa.tilelang_kernel import fp8_index from sglang.srt.layers.attention.dsa.tilelang_kernel import fp8_index
page_size = forward_batch.token_to_kv_pool.page_size page_size = get_token_to_kv_pool().page_size
assert page_size == 64, "only support page size 64" assert page_size == 64, "only support page size 64"
assert len(weights.shape) == 3 assert len(weights.shape) == 3
@@ -1011,7 +1016,7 @@ class Indexer(MultiPlatformOp):
topk_indices_list = [] topk_indices_list = []
block_tables = forward_batch.req_to_token_pool.req_to_token[ block_tables = get_req_to_token_pool().req_to_token[
forward_batch.req_pool_indices, : forward_batch.req_pool_indices, :
] ]
strided_indices = torch.arange( strided_indices = torch.arange(
@@ -1036,12 +1041,12 @@ class Indexer(MultiPlatformOp):
weights_partial = weights[q_len_start:q_len_end] weights_partial = weights[q_len_start:q_len_end]
weights_partial = weights_partial.squeeze(-1).unsqueeze(0).contiguous() weights_partial = weights_partial.squeeze(-1).unsqueeze(0).contiguous()
k_fp8 = forward_batch.token_to_kv_pool.get_index_k_continuous( k_fp8 = get_token_to_kv_pool().get_index_k_continuous(
layer_id, layer_id,
seq_len, seq_len,
block_tables[i], block_tables[i],
) )
k_scale = forward_batch.token_to_kv_pool.get_index_k_scale_continuous( k_scale = get_token_to_kv_pool().get_index_k_scale_continuous(
layer_id, layer_id,
seq_len, seq_len,
block_tables[i], block_tables[i],
@@ -1093,18 +1098,18 @@ class Indexer(MultiPlatformOp):
and can_use_dsa_fused_store( and can_use_dsa_fused_store(
key.dtype, key.dtype,
forward_batch.out_cache_loc.dtype, forward_batch.out_cache_loc.dtype,
forward_batch.token_to_kv_pool.page_size, get_token_to_kv_pool().page_size,
) )
): ):
# NOTE: wrapper already normalizes shape/contiguity and asserts dtypes. # NOTE: wrapper already normalizes shape/contiguity and asserts dtypes.
buf = forward_batch.token_to_kv_pool.get_index_k_with_scale_buffer( buf = get_token_to_kv_pool().get_index_k_with_scale_buffer(
layer_id=layer_id layer_id=layer_id
) )
fused_store_index_k_cache( fused_store_index_k_cache(
key, key,
buf, buf,
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
forward_batch.token_to_kv_pool.page_size, get_token_to_kv_pool().page_size,
) )
return return
@@ -1114,8 +1119,8 @@ class Indexer(MultiPlatformOp):
# layout with page_size=1; the same kv_cache.view works for both cases # layout with page_size=1; the same kv_cache.view works for both cases
# because page_size is 1 there. # because page_size is 1 there.
if _use_aiter: if _use_aiter:
page_size = forward_batch.token_to_kv_pool.page_size page_size = get_token_to_kv_pool().page_size
buf = forward_batch.token_to_kv_pool.get_index_k_with_scale_buffer( buf = get_token_to_kv_pool().get_index_k_with_scale_buffer(
layer_id=layer_id layer_id=layer_id
) )
kv_cache = buf.view(-1, page_size, 132).view(fp8_dtype) kv_cache = buf.view(-1, page_size, 132).view(fp8_dtype)
@@ -1140,7 +1145,7 @@ class Indexer(MultiPlatformOp):
if not out_loc.is_contiguous(): if not out_loc.is_contiguous():
out_loc = out_loc.contiguous() out_loc = out_loc.contiguous()
forward_batch.token_to_kv_pool.set_index_k_scale_buffer( get_token_to_kv_pool().set_index_k_scale_buffer(
layer_id=layer_id, layer_id=layer_id,
loc=out_loc, loc=out_loc,
index_k=k_fp8, index_k=k_fp8,
@@ -1175,15 +1180,13 @@ class Indexer(MultiPlatformOp):
from sglang.srt.layers.attention.dsa.triton_kernel import act_quant from sglang.srt.layers.attention.dsa.triton_kernel import act_quant
if TYPE_CHECKING: if TYPE_CHECKING:
assert isinstance(forward_batch.token_to_kv_pool, DSATokenToKVPool) assert isinstance(get_token_to_kv_pool(), DSATokenToKVPool)
# When upstream uses fused FP8 RMSNorm+quant, activations may be passed as # When upstream uses fused FP8 RMSNorm+quant, activations may be passed as
# a tuple like (x_fp8, x_scale[, y]). Use `x_meta` for shape/device queries. # a tuple like (x_fp8, x_scale[, y]). Use `x_meta` for shape/device queries.
x_meta = x[0] if isinstance(x, tuple) else x x_meta = x[0] if isinstance(x, tuple) else x
metadata = forward_batch.attn_backend.get_indexer_metadata( metadata = get_attn_backend().get_indexer_metadata(layer_id, forward_batch)
layer_id, forward_batch
)
enable_dual_stream = ( enable_dual_stream = (
self.alt_stream is not None self.alt_stream is not None
@@ -1405,12 +1408,10 @@ class Indexer(MultiPlatformOp):
layer_scatter_modes=None, layer_scatter_modes=None,
dynamic_scale: torch.Tensor = None, dynamic_scale: torch.Tensor = None,
) -> torch.Tensor: ) -> torch.Tensor:
if forward_batch.attn_backend.forward_metadata.seq_lens_cpu_int is None: if get_attn_backend().forward_metadata.seq_lens_cpu_int is None:
actual_seq_lengths_kv = forward_batch.attn_backend.forward_metadata.seq_lens actual_seq_lengths_kv = get_attn_backend().forward_metadata.seq_lens
else: else:
actual_seq_lengths_kv = ( actual_seq_lengths_kv = get_attn_backend().forward_metadata.seq_lens_cpu_int
forward_batch.attn_backend.forward_metadata.seq_lens_cpu_int
)
is_prefill = ( is_prefill = (
forward_batch.forward_mode.is_extend() forward_batch.forward_mode.is_extend()
and not forward_batch.forward_mode.is_draft_extend_v2() and not forward_batch.forward_mode.is_draft_extend_v2()
@@ -1558,7 +1559,7 @@ class Indexer(MultiPlatformOp):
torch.npu.current_stream(), torch.npu.current_stream(),
) )
forward_batch.token_to_kv_pool.set_index_k_buffer( get_token_to_kv_pool().set_index_k_buffer(
layer_id, forward_batch.out_cache_loc, k layer_id, forward_batch.out_cache_loc, k
) )
if is_prefill: if is_prefill:
@@ -1566,7 +1567,7 @@ class Indexer(MultiPlatformOp):
self.dsa_enable_prefill_cp self.dsa_enable_prefill_cp
and forward_batch.attn_cp_metadata is not None and forward_batch.attn_cp_metadata is not None
): ):
forward_batch.attn_backend.forward_metadata.actual_seq_lengths_q = ( get_attn_backend().forward_metadata.actual_seq_lengths_q = (
forward_batch.attn_cp_metadata.actual_seq_q_prev_tensor, forward_batch.attn_cp_metadata.actual_seq_q_prev_tensor,
forward_batch.attn_cp_metadata.actual_seq_q_next_tensor, forward_batch.attn_cp_metadata.actual_seq_q_next_tensor,
) )
@@ -1579,34 +1580,32 @@ class Indexer(MultiPlatformOp):
forward_batch.attn_cp_metadata.kv_len_next_tensor forward_batch.attn_cp_metadata.kv_len_next_tensor
+ forward_batch.extend_prefix_lens.squeeze() + forward_batch.extend_prefix_lens.squeeze()
) )
forward_batch.attn_backend.forward_metadata.actual_seq_lengths_kv = ( get_attn_backend().forward_metadata.actual_seq_lengths_kv = (
total_kv_len_prev_tensor, total_kv_len_prev_tensor,
total_kv_len_next_tensor, total_kv_len_next_tensor,
) )
else: else:
forward_batch.attn_backend.forward_metadata.actual_seq_lengths_kv = ( get_attn_backend().forward_metadata.actual_seq_lengths_kv = (
forward_batch.attn_cp_metadata.kv_len_prev_tensor, forward_batch.attn_cp_metadata.kv_len_prev_tensor,
forward_batch.attn_cp_metadata.kv_len_next_tensor, forward_batch.attn_cp_metadata.kv_len_next_tensor,
) )
actual_seq_lengths_q = ( actual_seq_lengths_q = (
forward_batch.attn_backend.forward_metadata.actual_seq_lengths_q get_attn_backend().forward_metadata.actual_seq_lengths_q
) )
actual_seq_lengths_kv = ( actual_seq_lengths_kv = (
forward_batch.attn_backend.forward_metadata.actual_seq_lengths_kv get_attn_backend().forward_metadata.actual_seq_lengths_kv
) )
else: else:
actual_seq_lengths_kv = forward_batch.seq_lens actual_seq_lengths_kv = forward_batch.seq_lens
actual_seq_lengths_q = forward_batch.extend_seq_lens.cumsum(dim=0) actual_seq_lengths_q = forward_batch.extend_seq_lens.cumsum(dim=0)
else: else:
if forward_batch.attn_backend.forward_metadata.actual_seq_lengths_q is None: if get_attn_backend().forward_metadata.actual_seq_lengths_q is None:
if ( if (
forward_batch.forward_mode.is_draft_extend_v2() forward_batch.forward_mode.is_draft_extend_v2()
or forward_batch.forward_mode.is_target_verify() or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend() or forward_batch.forward_mode.is_draft_extend()
): ):
num_draft_tokens = ( num_draft_tokens = get_attn_backend().speculative_num_draft_tokens
forward_batch.attn_backend.speculative_num_draft_tokens
)
actual_seq_lengths_q = torch.arange( actual_seq_lengths_q = torch.arange(
num_draft_tokens, num_draft_tokens,
num_draft_tokens + bs, num_draft_tokens + bs,
@@ -1622,10 +1621,10 @@ class Indexer(MultiPlatformOp):
) )
else: else:
actual_seq_lengths_q = ( actual_seq_lengths_q = (
forward_batch.attn_backend.forward_metadata.actual_seq_lengths_q get_attn_backend().forward_metadata.actual_seq_lengths_q
) )
past_key_states = forward_batch.token_to_kv_pool.get_index_k_buffer(layer_id) past_key_states = get_token_to_kv_pool().get_index_k_buffer(layer_id)
if self.rotary_emb.is_neox_style and self.alt_stream is not None: if self.rotary_emb.is_neox_style and self.alt_stream is not None:
torch.npu.current_stream().wait_event(q_rope_event) torch.npu.current_stream().wait_event(q_rope_event)
@@ -1637,7 +1636,7 @@ class Indexer(MultiPlatformOp):
and layer_scatter_modes.attn_mode == ScatterMode.TP_ATTN_FULL and layer_scatter_modes.attn_mode == ScatterMode.TP_ATTN_FULL
): ):
weights = scattered_to_tp_attn_full(weights, forward_batch) weights = scattered_to_tp_attn_full(weights, forward_batch)
block_table = forward_batch.attn_backend.forward_metadata.block_tables block_table = get_attn_backend().forward_metadata.block_tables
if ( if (
is_prefill is_prefill
and self.dsa_enable_prefill_cp and self.dsa_enable_prefill_cp
@@ -330,6 +330,14 @@ class DeepseekSparseAttnBackend(
self.qk_rope_head_dim = model_runner.model_config.qk_rope_head_dim self.qk_rope_head_dim = model_runner.model_config.qk_rope_head_dim
assert model_runner.req_to_token_pool is not None assert model_runner.req_to_token_pool is not None
# Pool refs — captured at construction so they survive deletion of the
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
# Keep a runner ref to read live state set after backend construction
# (e.g. hisparse_coordinator is built in model_runner *after*
# init_attention_backend()).
self.model_runner = model_runner
self.req_to_token = model_runner.req_to_token_pool.req_to_token self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.use_mha: bool = False self.use_mha: bool = False
@@ -392,6 +400,12 @@ class DeepseekSparseAttnBackend(
else: else:
self.workspace_buffer = None self.workspace_buffer = None
@property
def hisparse_coordinator(self):
# Live read: model_runner builds the coordinator *after*
# init_attention_backend(), so we cannot capture at __init__ time.
return self.model_runner.hisparse_coordinator
def get_device_int32_arange(self, l: int) -> torch.Tensor: def get_device_int32_arange(self, l: int) -> torch.Tensor:
if l > len(self._arange_buf): if l > len(self._arange_buf):
next_pow_of_2 = 1 << (l - 1).bit_length() next_pow_of_2 = 1 << (l - 1).bit_length()
@@ -425,7 +439,7 @@ class DeepseekSparseAttnBackend(
assert forward_batch.seq_lens_cpu is not None assert forward_batch.seq_lens_cpu is not None
max_seqlen_k = int(forward_batch.seq_lens_cpu.max().item() + draft_token_num) max_seqlen_k = int(forward_batch.seq_lens_cpu.max().item() + draft_token_num)
# [b, max_seqlen_k] # [b, max_seqlen_k]
page_table = forward_batch.req_to_token_pool.req_to_token[ page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, :max_seqlen_k forward_batch.req_pool_indices, :max_seqlen_k
] ]
@@ -580,8 +594,7 @@ class DeepseekSparseAttnBackend(
# Check if MHA FP8 dequantization is needed # Check if MHA FP8 dequantization is needed
mha_dequantize_needed = ( mha_dequantize_needed = (
self.use_mha self.use_mha and self.token_to_kv_pool.dtype == torch.float8_e4m3fn
and forward_batch.token_to_kv_pool.dtype == torch.float8_e4m3fn
) )
forward_batch.using_mha_one_shot_fp8_dequant = mha_dequantize_needed forward_batch.using_mha_one_shot_fp8_dequant = mha_dequantize_needed
@@ -606,8 +619,7 @@ class DeepseekSparseAttnBackend(
# Validate indices when logical tokens exceed physical capacity # Validate indices when logical tokens exceed physical capacity
# This is likely to be triggered by PP with high kv reuse & parallelism # This is likely to be triggered by PP with high kv reuse & parallelism
kv_cache_capacity = ( kv_cache_capacity = (
forward_batch.token_to_kv_pool.size self.token_to_kv_pool.size + self.token_to_kv_pool.page_size
+ forward_batch.token_to_kv_pool.page_size
) )
if forward_batch.seq_lens_sum > kv_cache_capacity: if forward_batch.seq_lens_sum > kv_cache_capacity:
max_idx = page_table_1_flattened.max().item() max_idx = page_table_1_flattened.max().item()
@@ -1380,7 +1392,7 @@ class DeepseekSparseAttnBackend(
if not layer.is_cross_attention if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
forward_batch.token_to_kv_pool.set_mla_kv_buffer( # type: ignore self.token_to_kv_pool.set_mla_kv_buffer( # type: ignore
layer, layer,
cache_loc, cache_loc,
k, k,
@@ -1405,7 +1417,7 @@ class DeepseekSparseAttnBackend(
# Do absorbed multi-latent attention (MLA path) # Do absorbed multi-latent attention (MLA path)
assert q_rope is not None assert q_rope is not None
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
if q_rope is not None: if q_rope is not None:
q_nope = q.view(-1, layer.tp_q_head_num, layer.v_head_dim) q_nope = q.view(-1, layer.tp_q_head_num, layer.v_head_dim)
@@ -1451,11 +1463,9 @@ class DeepseekSparseAttnBackend(
) )
# todo hisparse: to cover more backends # todo hisparse: to cover more backends
if forward_batch.hisparse_coordinator is not None: if self.hisparse_coordinator is not None:
page_table_1 = ( page_table_1 = self.token_to_kv_pool.translate_loc_to_hisparse_device(
forward_batch.token_to_kv_pool.translate_loc_to_hisparse_device( page_table_1
page_table_1
)
) )
if dsa_impl == "tilelang": if dsa_impl == "tilelang":
@@ -1580,7 +1590,7 @@ class DeepseekSparseAttnBackend(
if not layer.is_cross_attention if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
forward_batch.token_to_kv_pool.set_mla_kv_buffer( # type: ignore self.token_to_kv_pool.set_mla_kv_buffer( # type: ignore
layer, layer,
cache_loc, cache_loc,
k, k,
@@ -1588,7 +1598,7 @@ class DeepseekSparseAttnBackend(
) )
# Do absorbed multi-latent attention # Do absorbed multi-latent attention
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
if q_rope is not None: if q_rope is not None:
q_nope = q.view(-1, layer.tp_q_head_num, layer.v_head_dim) q_nope = q.view(-1, layer.tp_q_head_num, layer.v_head_dim)
q_rope = q_rope.view( q_rope = q_rope.view(
@@ -1609,8 +1619,8 @@ class DeepseekSparseAttnBackend(
if topk_indices is not None: if topk_indices is not None:
topk_indices = self._pad_topk_indices(topk_indices, q_nope.shape[0]) topk_indices = self._pad_topk_indices(topk_indices, q_nope.shape[0])
if forward_batch.hisparse_coordinator is not None: if self.hisparse_coordinator is not None:
page_table_1 = forward_batch.hisparse_coordinator.swap_in_selected_pages( page_table_1 = self.hisparse_coordinator.swap_in_selected_pages(
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.seq_lens, forward_batch.seq_lens,
topk_indices, topk_indices,
@@ -2105,11 +2115,9 @@ class DeepseekSparseAttnBackend(
if not layer.is_cross_attention if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.token_to_kv_pool.set_mla_kv_buffer(layer, cache_loc, k, k_rope)
layer, cache_loc, k, k_rope
)
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
kv_cache = k_cache.view(-1, self.real_page_size, self.kv_cache_dim).unsqueeze(1) kv_cache = k_cache.view(-1, self.real_page_size, self.kv_cache_dim).unsqueeze(1)
if merge_query: if merge_query:
@@ -2221,12 +2229,11 @@ class DeepseekSparseAttnBackend(
) # SM90/SM100 only ) # SM90/SM100 only
and max_kv_len and max_kv_len
<= envs.SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD.get() # Short enough for MHA <= envs.SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD.get() # Short enough for MHA
and forward_batch.token_to_kv_pool.dtype and self.token_to_kv_pool.dtype in [torch.bfloat16, torch.float8_e4m3fn]
in [torch.bfloat16, torch.float8_e4m3fn]
and sum_seq_lens and sum_seq_lens
<= forward_batch.get_max_chunk_capacity() # Fits in chunk <= forward_batch.get_max_chunk_capacity() # Fits in chunk
and (not is_dsa_enable_prefill_cp()) # CP not enabled and (not is_dsa_enable_prefill_cp()) # CP not enabled
and (forward_batch.hisparse_coordinator is None) and (self.hisparse_coordinator is None)
) )
else: else:
self.use_mha = False # Decode/verify always use MLA self.use_mha = False # Decode/verify always use MLA
@@ -2272,7 +2279,7 @@ class DeepseekSparseAttnBackend(
self, layer_id: int, forward_batch: ForwardBatch self, layer_id: int, forward_batch: ForwardBatch
) -> DSAIndexerMetadata: ) -> DSAIndexerMetadata:
force_unfused = ( force_unfused = (
forward_batch.hisparse_coordinator is not None self.hisparse_coordinator is not None
and forward_batch.forward_mode.is_decode_or_idle() and forward_batch.forward_mode.is_decode_or_idle()
) )
return DSAIndexerMetadata( return DSAIndexerMetadata(
@@ -23,6 +23,7 @@ from sglang.srt.mem_cache.deepseek_v4_compress_state import (
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import ( from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import (
DeepseekV4HipRadixBackend, DeepseekV4HipRadixBackend,
) )
@@ -99,24 +100,26 @@ class CompressorHip(_CompressorBase):
def use_hip_fused_compress(self) -> bool: def use_hip_fused_compress(self) -> bool:
return envs.SGLANG_OPT_USE_FUSED_COMPRESS.get() return envs.SGLANG_OPT_USE_FUSED_COMPRESS.get()
def _get_states(self, forward_batch: ForwardBatch) -> KVAndScore: def _get_states(
token_to_kv_pool = forward_batch.token_to_kv_pool self,
forward_batch: ForwardBatch,
attn_backend: AttentionBackend,
) -> KVAndScore:
token_to_kv_pool = attn_backend.token_to_kv_pool
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
if self.is_in_indexer: if self.is_in_indexer:
return token_to_kv_pool.get_indexer_compress_states(self.layer_id) return token_to_kv_pool.get_indexer_compress_states(self.layer_id)
else: else:
return token_to_kv_pool.get_attention_compress_states(self.layer_id) return token_to_kv_pool.get_attention_compress_states(self.layer_id)
def _get_state_pool(self, forward_batch: ForwardBatch) -> CompressStatePool: def _get_state_pool(self, attn_backend: AttentionBackend) -> CompressStatePool:
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = attn_backend.token_to_kv_pool
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
if self.is_in_indexer: if self.is_in_indexer:
ret = token_to_kv_pool.get_indexer_compress_states(self.layer_id) ret = token_to_kv_pool.get_indexer_compress_states(self.layer_id)
else: else:
ret = token_to_kv_pool.get_attention_compress_states(self.layer_id) ret = token_to_kv_pool.get_attention_compress_states(self.layer_id)
assert isinstance(ret, CompressStatePool) assert isinstance(ret, CompressStatePool)
return ret return ret
def overlap_transform(self, tensor: torch.Tensor, fill_value: Any) -> torch.Tensor: def overlap_transform(self, tensor: torch.Tensor, fill_value: Any) -> torch.Tensor:
@@ -155,18 +158,19 @@ class CompressorHip(_CompressorBase):
self, self,
kv_and_scores: KVAndScore, kv_and_scores: KVAndScore,
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
attn_backend: AttentionBackend,
): ):
backend = forward_batch.attn_backend backend = attn_backend
if TYPE_CHECKING: if TYPE_CHECKING:
assert isinstance(backend, DeepseekV4HipRadixBackend) assert isinstance(backend, DeepseekV4HipRadixBackend)
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = backend.token_to_kv_pool
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
state_pool = self._get_state_pool(forward_batch) state_pool = self._get_state_pool(backend)
prefix_lens = forward_batch.extend_prefix_lens_cpu prefix_lens = forward_batch.extend_prefix_lens_cpu
extend_lens = forward_batch.extend_seq_lens_cpu extend_lens = forward_batch.extend_seq_lens_cpu
req_pool_indices = forward_batch.req_pool_indices req_pool_indices = forward_batch.req_pool_indices
req_to_token = forward_batch.req_to_token_pool.req_to_token req_to_token = backend.req_to_token_pool.req_to_token
assert not self.forward_mode.is_target_verify() assert not self.forward_mode.is_target_verify()
assert extend_lens is not None and prefix_lens is not None assert extend_lens is not None and prefix_lens is not None
@@ -289,18 +293,19 @@ class CompressorHip(_CompressorBase):
self, self,
kv_and_scores: KVAndScore, kv_and_scores: KVAndScore,
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
attn_backend: AttentionBackend,
): ):
"""Paged and cudagraph compatible version of compress_decode""" """Paged and cudagraph compatible version of compress_decode"""
assert self.ape_converted assert self.ape_converted
state_pool = self._get_state_pool(forward_batch) state_pool = self._get_state_pool(attn_backend)
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = attn_backend.token_to_kv_pool
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
req_pool_indices = forward_batch.req_pool_indices req_pool_indices = forward_batch.req_pool_indices
req_to_token = forward_batch.req_to_token_pool.req_to_token req_to_token = attn_backend.req_to_token_pool.req_to_token
seq_lens = forward_batch.seq_lens seq_lens = forward_batch.seq_lens
if forward_batch.forward_mode.is_target_verify(): if forward_batch.forward_mode.is_target_verify():
draft_tokens = forward_batch.attn_backend.speculative_num_draft_tokens draft_tokens = attn_backend.speculative_num_draft_tokens
offsets = torch.arange(1, draft_tokens + 1, device=seq_lens.device) offsets = torch.arange(1, draft_tokens + 1, device=seq_lens.device)
seq_lens_2d = seq_lens[:, None] + offsets[None, :] seq_lens_2d = seq_lens[:, None] + offsets[None, :]
seq_lens = seq_lens_2d.view(-1) seq_lens = seq_lens_2d.view(-1)
@@ -378,11 +383,12 @@ class CompressorHip(_CompressorBase):
self, self,
kv_score: torch.Tensor, kv_score: torch.Tensor,
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
attn_backend: AttentionBackend,
) -> torch.Tensor: ) -> torch.Tensor:
backend = forward_batch.attn_backend backend = attn_backend
if TYPE_CHECKING: if TYPE_CHECKING:
assert isinstance(backend, DeepseekV4HipRadixBackend) assert isinstance(backend, DeepseekV4HipRadixBackend)
kv_score_buffer = self._get_state_pool(forward_batch) kv_score_buffer = self._get_state_pool(backend)
kv_score_buffer = kv_score_buffer.kv_score_buffer.kv_score kv_score_buffer = kv_score_buffer.kv_score_buffer.kv_score
return backend.forward_compress( return backend.forward_compress(
@@ -402,9 +408,12 @@ class CompressorHip(_CompressorBase):
self, self,
kv_score: torch.Tensor, kv_score: torch.Tensor,
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
attn_backend: AttentionBackend,
) -> torch.Tensor: ) -> torch.Tensor:
if self.use_fused_compress: if self.use_fused_compress:
return self.compress_fused(kv_score, forward_batch) return self.compress_fused(
kv_score, forward_batch, attn_backend=attn_backend
)
self.compress_decode = self.compress_decode_paged self.compress_decode = self.compress_decode_paged
self.compress_extend = self.compress_extend_paged self.compress_extend = self.compress_extend_paged
@@ -420,11 +429,13 @@ class CompressorHip(_CompressorBase):
result = self.compress_decode( result = self.compress_decode(
kv_and_scores=kv_and_scores, kv_and_scores=kv_and_scores,
forward_batch=forward_batch, forward_batch=forward_batch,
attn_backend=attn_backend,
) )
elif forward_batch.forward_mode.is_extend(): elif forward_batch.forward_mode.is_extend():
result = self.compress_extend( result = self.compress_extend(
kv_and_scores=kv_and_scores, kv_and_scores=kv_and_scores,
forward_batch=forward_batch, forward_batch=forward_batch,
attn_backend=attn_backend,
) )
else: else:
msg = f"Forward mode {forward_batch.forward_mode} not supported in Compressor." msg = f"Forward mode {forward_batch.forward_mode} not supported in Compressor."
@@ -445,11 +456,17 @@ class CompressorHip(_CompressorBase):
setattr(forward_batch, attr, decoded) setattr(forward_batch, attr, decoded)
return decoded return decoded
def forward(self, x: torch.Tensor, forward_batch: ForwardBatch) -> torch.Tensor: def forward(
self,
x: torch.Tensor,
forward_batch: ForwardBatch,
attn_backend: AttentionBackend,
) -> torch.Tensor:
if forward_batch.forward_mode.is_idle(): if forward_batch.forward_mode.is_idle():
assert x.shape[0] == 0 assert x.shape[0] == 0
return x.new_empty(0, self.head_dim) return x.new_empty(0, self.head_dim)
kv_score = self.compute_kv_score(x, forward_batch) kv_score = self.compute_kv_score(x, forward_batch)
self.forward_mode = forward_batch.forward_mode self.forward_mode = forward_batch.forward_mode
return self.compress_dispatch(kv_score, forward_batch) return self.compress_dispatch(
kv_score, forward_batch, attn_backend=attn_backend
)
@@ -31,6 +31,7 @@ from sglang.srt.models.deepseek_v2 import _is_hip
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.deepseek_v4_backend import DeepseekV4AttnBackend from sglang.srt.layers.attention.deepseek_v4_backend import DeepseekV4AttnBackend
from sglang.srt.layers.rotary_embedding import RotaryEmbedding from sglang.srt.layers.rotary_embedding import RotaryEmbedding
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@@ -123,11 +124,11 @@ class CompressorBackendMixin:
# attn_backend.forward(), so Raw -> DSV4Metadata must happen here too # attn_backend.forward(), so Raw -> DSV4Metadata must happen here too
# (e.g. 1.6T layer 0 has compress_ratio=128 and needs cX_compress_metadata). # (e.g. 1.6T layer 0 has compress_ratio=128 and needs cX_compress_metadata).
self._maybe_upgrade_forward_metadata() self._maybe_upgrade_forward_metadata()
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = self.token_to_kv_pool
if TYPE_CHECKING: if TYPE_CHECKING:
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
new_compressed_kv = compressor(x, forward_batch) new_compressed_kv = compressor(x, forward_batch, attn_backend=self)
core_metadata = self.forward_metadata.core_metadata core_metadata = self.forward_metadata.core_metadata
out_loc = ( out_loc = (
core_metadata.c4_out_loc core_metadata.c4_out_loc
@@ -154,11 +155,11 @@ class CompressorBackendMixin:
assert is_overlap_compress(compressor.ratio) assert is_overlap_compress(compressor.ratio)
# PREP_IN_CG lazy upgrade (see forward_core_compressor for rationale). # PREP_IN_CG lazy upgrade (see forward_core_compressor for rationale).
self._maybe_upgrade_forward_metadata() self._maybe_upgrade_forward_metadata()
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = self.token_to_kv_pool
if TYPE_CHECKING: if TYPE_CHECKING:
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
new_compressed_kv = compressor(x, forward_batch) new_compressed_kv = compressor(x, forward_batch, attn_backend=self)
if envs.SGLANG_OPT_USE_FUSED_STORE_CACHE.get(): if envs.SGLANG_OPT_USE_FUSED_STORE_CACHE.get():
token_to_kv_pool.set_index_k_fused( token_to_kv_pool.set_index_k_fused(
layer_id=layer_id, layer_id=layer_id,
@@ -339,20 +340,16 @@ class Compressor(nn.Module):
ape = torch.cat([ape[0], ape[1]], dim=0) ape = torch.cat([ape[0], ape[1]], dim=0)
self.ape.data.copy_(ape.view(self.ratio, -1)) self.ape.data.copy_(ape.view(self.ratio, -1))
# NOTE: used by v2 compressor backend def get_state_pool(self, attn_backend: AttentionBackend) -> CompressStatePool:
def get_state_pool(self, forward_batch: ForwardBatch) -> CompressStatePool: token_to_kv_pool = attn_backend.token_to_kv_pool
token_to_kv_pool = forward_batch.token_to_kv_pool
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
if self.is_in_indexer: if self.is_in_indexer:
ret = token_to_kv_pool.get_indexer_compress_states(self.layer_id) ret = token_to_kv_pool.get_indexer_compress_states(self.layer_id)
else: else:
ret = token_to_kv_pool.get_attention_compress_states(self.layer_id) ret = token_to_kv_pool.get_attention_compress_states(self.layer_id)
assert isinstance(ret, CompressStatePool) assert isinstance(ret, CompressStatePool)
return ret return ret
# NOTE: used by v2 compressor backend
def compute_kv_score(self, x: torch.Tensor, forward_batch: ForwardBatch): def compute_kv_score(self, x: torch.Tensor, forward_batch: ForwardBatch):
kv_score = linear_bf16_fp32(x, self.wkv_gate.weight) kv_score = linear_bf16_fp32(x, self.wkv_gate.weight)
@@ -366,19 +363,22 @@ class Compressor(nn.Module):
) )
return kv_score return kv_score
def forward(self, x: torch.Tensor, forward_batch: ForwardBatch) -> torch.Tensor: def forward(
self,
x: torch.Tensor,
forward_batch: ForwardBatch,
attn_backend: AttentionBackend,
) -> torch.Tensor:
if forward_batch.forward_mode.is_idle(): if forward_batch.forward_mode.is_idle():
assert x.shape[0] == 0 assert x.shape[0] == 0
return x.new_empty(0, self.head_dim) return x.new_empty(0, self.head_dim)
kv_score = self.compute_kv_score(x, forward_batch) kv_score = self.compute_kv_score(x, forward_batch)
backend = forward_batch.attn_backend
if TYPE_CHECKING: if TYPE_CHECKING:
assert isinstance(backend, DeepseekV4AttnBackend) assert isinstance(attn_backend, DeepseekV4AttnBackend)
kv_score_buffer = self.get_state_pool(forward_batch) kv_score_buffer = self.get_state_pool(attn_backend).kv_score_buffer.kv_score
kv_score_buffer = kv_score_buffer.kv_score_buffer.kv_score return attn_backend.forward_compress(
return backend.forward_compress(
kv_score_buffer=kv_score_buffer, kv_score_buffer=kv_score_buffer,
kv_score_input=kv_score, kv_score_input=kv_score,
ape=self.ape.view(-1, self.head_dim), ape=self.ape.view(-1, self.head_dim),
@@ -106,10 +106,10 @@ class CompressorBackendMixin:
return return
self._maybe_upgrade_forward_metadata() self._maybe_upgrade_forward_metadata()
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = self.token_to_kv_pool
token_to_kv_pool = cast("DeepSeekV4TokenToKVPool", token_to_kv_pool) token_to_kv_pool = cast("DeepSeekV4TokenToKVPool", token_to_kv_pool)
kv_score_input = compressor.compute_kv_score(x, forward_batch) kv_score_input = compressor.compute_kv_score(x, forward_batch)
state_pool = compressor.get_state_pool(forward_batch) state_pool = compressor.get_state_pool(self)
out_loc = self._get_out_loc(compressor.ratio) out_loc = self._get_out_loc(compressor.ratio)
if compressor.is_in_indexer: if compressor.is_in_indexer:
kv_cache = token_to_kv_pool.get_index_k_with_scale_buffer(layer_id) kv_cache = token_to_kv_pool.get_index_k_with_scale_buffer(layer_id)
@@ -22,6 +22,7 @@ from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer
from sglang.srt.utils import add_prefix, is_hip from sglang.srt.utils import add_prefix, is_hip
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.dsv4.compressor import ( from sglang.srt.layers.attention.dsv4.compressor import (
CompressorBackendMixin, CompressorBackendMixin,
) )
@@ -321,7 +322,7 @@ class C4IndexerBackendMixin:
# PREP_IN_CG lazy upgrade: this runs from MQALayer._forward_prepare, # PREP_IN_CG lazy upgrade: this runs from MQALayer._forward_prepare,
# before attn_backend.forward() would trigger the upgrade. # before attn_backend.forward() would trigger the upgrade.
self._maybe_upgrade_forward_metadata() self._maybe_upgrade_forward_metadata()
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = self.token_to_kv_pool
if TYPE_CHECKING: if TYPE_CHECKING:
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
@@ -398,7 +399,7 @@ class C4IndexerBackendMixin:
indexer_capturer = get_global_indexer_capturer() indexer_capturer = get_global_indexer_capturer()
capture_enabled = indexer_capturer is not None capture_enabled = indexer_capturer is not None
hisparse_coordinator = forward_batch.hisparse_coordinator hisparse_coordinator = self.hisparse_coordinator
hisparse_decode = ( hisparse_decode = (
hisparse_coordinator is not None and forward_batch.forward_mode.is_decode() hisparse_coordinator is not None and forward_batch.forward_mode.is_decode()
) )
@@ -541,10 +542,11 @@ class C4Indexer(nn.Module):
x: torch.Tensor, x: torch.Tensor,
q_lora: torch.Tensor, q_lora: torch.Tensor,
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
attn_backend: AttentionBackend,
enable_multi_stream: bool = False, enable_multi_stream: bool = False,
q_lora_ready: Optional[torch.cuda.Event] = None, q_lora_ready: Optional[torch.cuda.Event] = None,
) -> None: ) -> None:
return forward_batch.attn_backend.forward_c4_indexer( return attn_backend.forward_c4_indexer(
x=x, x=x,
q_lora=q_lora, q_lora=q_lora,
forward_batch=forward_batch, forward_batch=forward_batch,
@@ -117,6 +117,10 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
) )
self.head_size = model_runner.model_config.head_dim self.head_size = model_runner.model_config.head_dim
# Pool refs — captured at construction so they survive deletion of the
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.req_to_token = model_runner.req_to_token_pool.req_to_token self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.kv_cache_dtype = model_runner.kv_cache_dtype self.kv_cache_dtype = model_runner.kv_cache_dtype
self.kv_cache_dtype_str = model_runner.server_args.kv_cache_dtype self.kv_cache_dtype_str = model_runner.server_args.kv_cache_dtype
@@ -183,7 +187,7 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
metadata.orig_seq_lens_tensor = forward_batch.orig_seq_lens metadata.orig_seq_lens_tensor = forward_batch.orig_seq_lens
metadata.orig_seq_lens = forward_batch.orig_seq_lens.tolist() metadata.orig_seq_lens = forward_batch.orig_seq_lens.tolist()
metadata.block_tables = forward_batch.req_to_token_pool.req_to_token[ metadata.block_tables = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len forward_batch.req_pool_indices, : metadata.max_seq_len
] ]
# Convert the block table to a strided format. # Convert the block table to a strided format.
@@ -346,9 +350,7 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
assert current_end <= self.max_context_len assert current_end <= self.max_context_len
# Do multi-head attention # Do multi-head attention
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer( key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
layer.layer_id
)
key_cache = key_cache.view( key_cache = key_cache.view(
-1, self.page_size, layer.tp_k_head_num, layer.head_dim -1, self.page_size, layer.tp_k_head_num, layer.head_dim
) )
@@ -358,7 +360,7 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
if key is not None and value is not None: if key is not None and value is not None:
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, layer,
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
key, key,
@@ -442,9 +444,7 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
key = k.view(-1, self.num_kv_heads, self.head_size) key = k.view(-1, self.num_kv_heads, self.head_size)
value = v.view(-1, self.num_kv_heads, self.head_size) value = v.view(-1, self.num_kv_heads, self.head_size)
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer( key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
layer.layer_id
)
key_cache = key_cache.view( key_cache = key_cache.view(
-1, self.page_size, layer.tp_k_head_num, layer.head_dim -1, self.page_size, layer.tp_k_head_num, layer.head_dim
) )
@@ -454,7 +454,7 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
if key is not None and value is not None: if key is not None and value is not None:
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, layer,
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
key, key,
@@ -126,6 +126,10 @@ class FlashAttentionBackend(AttentionBackend):
self.device = model_runner.device self.device = model_runner.device
self.decode_cuda_graph_metadata = {} self.decode_cuda_graph_metadata = {}
self.target_verify_metadata = {} self.target_verify_metadata = {}
# Pool refs — captured at construction so they survive deletion of the
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.req_to_token = model_runner.req_to_token_pool.req_to_token self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.kv_cache_dtype = model_runner.kv_cache_dtype self.kv_cache_dtype = model_runner.kv_cache_dtype
self.kv_cache_dtype_str = model_runner.server_args.kv_cache_dtype self.kv_cache_dtype_str = model_runner.server_args.kv_cache_dtype
@@ -138,8 +142,6 @@ class FlashAttentionBackend(AttentionBackend):
isinstance(model_runner.token_to_kv_pool, SWAKVPool) isinstance(model_runner.token_to_kv_pool, SWAKVPool)
and model_runner.token_to_kv_pool.swa_layer_nums > 0 and model_runner.token_to_kv_pool.swa_layer_nums > 0
) )
if self.use_sliding_window_kv_pool:
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.topk = model_runner.server_args.speculative_eagle_topk or 0
self.speculative_num_steps = speculative_num_steps self.speculative_num_steps = speculative_num_steps
@@ -295,7 +297,7 @@ class FlashAttentionBackend(AttentionBackend):
), ),
(1, 0), (1, 0),
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
else: else:
@@ -315,7 +317,7 @@ class FlashAttentionBackend(AttentionBackend):
), ),
(1, 0), (1, 0),
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
metadata_expand = FlashAttentionMetadata() metadata_expand = FlashAttentionMetadata()
@@ -358,7 +360,7 @@ class FlashAttentionBackend(AttentionBackend):
metadata.cu_seqlens_k = torch.nn.functional.pad( metadata.cu_seqlens_k = torch.nn.functional.pad(
torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0) torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0)
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
# Precompute FA3 scheduler metadata to avoid per-layer # Precompute FA3 scheduler metadata to avoid per-layer
@@ -394,7 +396,7 @@ class FlashAttentionBackend(AttentionBackend):
), ),
(1, 0), (1, 0),
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
@@ -416,7 +418,7 @@ class FlashAttentionBackend(AttentionBackend):
), ),
(1, 0), (1, 0),
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
@@ -477,7 +479,7 @@ class FlashAttentionBackend(AttentionBackend):
) )
_, sort_order = torch.sort(keys, dim=1) _, sort_order = torch.sort(keys, dim=1)
non_masked_page_table = ( non_masked_page_table = (
forward_batch.req_to_token_pool.req_to_token[ self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : forward_batch.req_pool_indices, :
] ]
.gather(1, cols) .gather(1, cols)
@@ -506,7 +508,7 @@ class FlashAttentionBackend(AttentionBackend):
metadata.cu_seqlens_k = torch.nn.functional.pad( metadata.cu_seqlens_k = torch.nn.functional.pad(
torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0) torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0)
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
@@ -538,12 +540,12 @@ class FlashAttentionBackend(AttentionBackend):
(1, 0), (1, 0),
) )
metadata.encoder_max_seq_len_k = metadata.encoder_lens_int32.max().item() metadata.encoder_max_seq_len_k = metadata.encoder_lens_int32.max().item()
metadata.encoder_page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.encoder_page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.encoder_max_seq_len_k forward_batch.req_pool_indices, : metadata.encoder_max_seq_len_k
] ]
# Currently only support forward_batch.encoder_lens.numel() == 1 # Currently only support forward_batch.encoder_lens.numel() == 1
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
metadata.encoder_max_seq_len_k : ( metadata.encoder_max_seq_len_k : (
metadata.encoder_max_seq_len_k + metadata.max_seq_len_k metadata.encoder_max_seq_len_k + metadata.max_seq_len_k
@@ -641,11 +643,11 @@ class FlashAttentionBackend(AttentionBackend):
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
if not self.use_mla: if not self.use_mla:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
else: else:
forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.token_to_kv_pool.set_mla_kv_buffer(
layer, layer,
cache_loc, cache_loc,
k, k,
@@ -746,9 +748,7 @@ class FlashAttentionBackend(AttentionBackend):
# Use Flash Attention for prefill # Use Flash Attention for prefill
if not self.use_mla: if not self.use_mla:
# Do multi-head attention # Do multi-head attention
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer( key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
layer.layer_id
)
key_cache = key_cache.view( key_cache = key_cache.view(
-1, self.page_size, layer.tp_k_head_num, layer.head_dim -1, self.page_size, layer.tp_k_head_num, layer.head_dim
@@ -950,9 +950,9 @@ class FlashAttentionBackend(AttentionBackend):
else: else:
assert self.fa_impl_ver == 3, "Only FA3 support here" assert self.fa_impl_ver == 3, "Only FA3 support here"
# Do absorbed multi-latent attention # Do absorbed multi-latent attention
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer( kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
layer.layer_id q.dtype
).to(q.dtype) )
k_rope = kv_cache[:, :, layer.v_head_dim :] k_rope = kv_cache[:, :, layer.v_head_dim :]
c_kv = kv_cache[:, :, : layer.v_head_dim] c_kv = kv_cache[:, :, : layer.v_head_dim]
k_rope_cache = k_rope.view( k_rope_cache = k_rope.view(
@@ -1050,11 +1050,11 @@ class FlashAttentionBackend(AttentionBackend):
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
if not self.use_mla: if not self.use_mla:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
else: else:
forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.token_to_kv_pool.set_mla_kv_buffer(
layer, layer,
cache_loc, cache_loc,
k, k,
@@ -1113,9 +1113,7 @@ class FlashAttentionBackend(AttentionBackend):
if not self.use_mla: if not self.use_mla:
# Do multi-head attention # Do multi-head attention
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer( key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
layer.layer_id
)
key_cache = key_cache.view( key_cache = key_cache.view(
-1, self.page_size, layer.tp_k_head_num, layer.head_dim -1, self.page_size, layer.tp_k_head_num, layer.head_dim
) )
@@ -1248,9 +1246,7 @@ class FlashAttentionBackend(AttentionBackend):
o = result o = result
else: else:
# Do absorbed multi-latent attention # Do absorbed multi-latent attention
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id).to( kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(q.dtype)
q.dtype
)
k_rope = kv_cache[:, :, layer.v_head_dim :] k_rope = kv_cache[:, :, layer.v_head_dim :]
c_kv = kv_cache[:, :, : layer.v_head_dim] c_kv = kv_cache[:, :, : layer.v_head_dim]
k_rope_cache = k_rope.view( k_rope_cache = k_rope.view(
@@ -126,7 +126,8 @@ class FlashInferAttnBackend(AttentionBackend):
self.prefill_backend = "fa2" self.prefill_backend = "fa2"
self.decode_backend = "fa2" self.decode_backend = "fa2"
# Store multi-item scoring flag for efficient access self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.enable_mis = model_runner.server_args.enable_mis self.enable_mis = model_runner.server_args.enable_mis
# FIXME: remove dllm workarounds from flashinfer # FIXME: remove dllm workarounds from flashinfer
@@ -802,7 +803,7 @@ class FlashInferAttnBackend(AttentionBackend):
if k is not None: if k is not None:
assert v is not None assert v is not None
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
@@ -812,7 +813,7 @@ class FlashInferAttnBackend(AttentionBackend):
) )
o = prefill_wrapper_paged.forward( o = prefill_wrapper_paged.forward(
q.view(-1, layer.tp_q_head_num, layer.head_dim), q.view(-1, layer.tp_q_head_num, layer.head_dim),
forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id), self.token_to_kv_pool.get_kv_buffer(layer.layer_id),
causal=causal, causal=causal,
sm_scale=layer.scaling, sm_scale=layer.scaling,
# Disable sliding window attention for multi-item scoring: # Disable sliding window attention for multi-item scoring:
@@ -836,12 +837,12 @@ class FlashInferAttnBackend(AttentionBackend):
) )
else: else:
# If `k`/`v` are not explicitly provided, fall back to the KV cache stored in # If `k`/`v` are not explicitly provided, fall back to the KV cache stored in
# `forward_batch.token_to_kv_pool` for this layer. This enables attention over # `self.token_to_kv_pool` for this layer. This enables attention over
# previously cached context without re-materializing KV tensors (e.g., the # previously cached context without re-materializing KV tensors (e.g., the
# IQuestLoopCoder path uses token_to_kv_pool as the KV source). # IQuestLoopCoder path uses token_to_kv_pool as the KV source).
if k is None and v is None: if k is None and v is None:
k = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id)[0] k = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)[0]
v = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id)[1] v = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)[1]
causal = True causal = True
if ( if (
layer.is_cross_attention layer.is_cross_attention
@@ -875,7 +876,7 @@ class FlashInferAttnBackend(AttentionBackend):
) )
o2, s2 = prefill_wrapper_paged.forward_return_lse( o2, s2 = prefill_wrapper_paged.forward_return_lse(
q.view(-1, layer.tp_q_head_num, layer.head_dim), q.view(-1, layer.tp_q_head_num, layer.head_dim),
forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id), self.token_to_kv_pool.get_kv_buffer(layer.layer_id),
causal=False, causal=False,
sm_scale=layer.scaling, sm_scale=layer.scaling,
logits_soft_cap=logits_soft_cap, logits_soft_cap=logits_soft_cap,
@@ -884,7 +885,7 @@ class FlashInferAttnBackend(AttentionBackend):
o, _ = merge_state(o1, s1, o2, s2) o, _ = merge_state(o1, s1, o2, s2)
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
@@ -912,14 +913,14 @@ class FlashInferAttnBackend(AttentionBackend):
if k is not None: if k is not None:
assert v is not None assert v is not None
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
# Call the wrapped function # Call the wrapped function
o = decode_wrapper.forward( o = decode_wrapper.forward(
q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim), q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim),
forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id), self.token_to_kv_pool.get_kv_buffer(layer.layer_id),
sm_scale=layer.scaling, sm_scale=layer.scaling,
logits_soft_cap=layer.logit_cap, logits_soft_cap=layer.logit_cap,
# Must use _float to avoid device-to-host copy that breaks cuda graph capture. # Must use _float to avoid device-to-host copy that breaks cuda graph capture.
@@ -1547,6 +1548,7 @@ class FlashInferMultiStepDraftBackend:
# Cached variables for generate_draft_decode_kv_indices # Cached variables for generate_draft_decode_kv_indices
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1] self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
self.req_to_token_pool = model_runner.req_to_token_pool
def common_template( def common_template(
self, self,
@@ -1562,7 +1564,7 @@ class FlashInferMultiStepDraftBackend:
(self.speculative_num_steps, num_seqs, self.topk) (self.speculative_num_steps, num_seqs, self.topk)
]( ](
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.req_to_token_pool.req_to_token, self.req_to_token_pool.req_to_token,
forward_batch.seq_lens, forward_batch.seq_lens,
kv_indices_buffer, kv_indices_buffer,
self.kv_indptr, self.kv_indptr,
@@ -204,6 +204,10 @@ class FlashInferMLAAttnBackend(AttentionBackend):
self.max_context_len = model_runner.model_config.context_len self.max_context_len = model_runner.model_config.context_len
self.device = model_runner.device self.device = model_runner.device
self.skip_prefill = skip_prefill self.skip_prefill = skip_prefill
# Pool refs — captured at construction so they survive deletion of the
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.enable_chunk_kv = ( self.enable_chunk_kv = (
not skip_prefill not skip_prefill
and get_global_server_args().disaggregation_mode != "decode" and get_global_server_args().disaggregation_mode != "decode"
@@ -544,11 +548,9 @@ class FlashInferMLAAttnBackend(AttentionBackend):
assert v is not None assert v is not None
if save_kv_cache: if save_kv_cache:
if k_rope is not None: if k_rope is not None:
forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.token_to_kv_pool.set_mla_kv_buffer(layer, cache_loc, k, k_rope)
layer, cache_loc, k, k_rope
)
else: else:
forward_batch.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
if q_rope is not None: if q_rope is not None:
q = q.view(-1, layer.tp_q_head_num, layer.v_head_dim) q = q.view(-1, layer.tp_q_head_num, layer.v_head_dim)
q_rope = q_rope.view( q_rope = q_rope.view(
@@ -572,9 +574,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
) )
else: else:
# mla paged prefill # mla paged prefill
k_buf = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id).to( k_buf = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(q.dtype)
q.dtype
)
if q_rope is None: if q_rope is None:
qall = q.view(-1, layer.tp_q_head_num, layer.head_dim) qall = q.view(-1, layer.tp_q_head_num, layer.head_dim)
q, q_rope = ( q, q_rope = (
@@ -611,14 +611,14 @@ class FlashInferMLAAttnBackend(AttentionBackend):
assert v is not None assert v is not None
if save_kv_cache: if save_kv_cache:
if k_rope is not None: if k_rope is not None:
forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.token_to_kv_pool.set_mla_kv_buffer(
layer, layer,
cache_loc, cache_loc,
k, k,
k_rope, k_rope,
) )
else: else:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, layer,
cache_loc, cache_loc,
k, k,
@@ -636,9 +636,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
q_nope = reshaped_q[:, :, : layer.v_head_dim] q_nope = reshaped_q[:, :, : layer.v_head_dim]
q_rope = reshaped_q[:, :, layer.v_head_dim :] q_rope = reshaped_q[:, :, layer.v_head_dim :]
k_buffer = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id).to( k_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(q.dtype)
q.dtype
)
o = q_nope.new_empty(q_nope.shape) o = q_nope.new_empty(q_nope.shape)
# Direct call to run without the wrapper # Direct call to run without the wrapper
@@ -944,6 +942,7 @@ class FlashInferMLAMultiStepDraftBackend:
self.max_context_len = self.attn_backends[0].max_context_len self.max_context_len = self.attn_backends[0].max_context_len
# Cached variables for generate_draft_decode_kv_indices # Cached variables for generate_draft_decode_kv_indices
self.req_to_token_pool = model_runner.req_to_token_pool
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1] self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
self.page_size = model_runner.server_args.page_size self.page_size = model_runner.server_args.page_size
@@ -961,7 +960,7 @@ class FlashInferMLAMultiStepDraftBackend:
(self.speculative_num_steps, num_seqs, self.topk) (self.speculative_num_steps, num_seqs, self.topk)
]( ](
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.req_to_token_pool.req_to_token, self.req_to_token_pool.req_to_token,
forward_batch.seq_lens, forward_batch.seq_lens,
kv_indices_buffer, kv_indices_buffer,
self.kv_indptr, self.kv_indptr,
@@ -410,14 +410,14 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
if k is not None: if k is not None:
assert v is not None assert v is not None
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, layer,
cache_loc, cache_loc,
k, k,
v, v,
) )
bs = forward_batch.batch_size bs = forward_batch.batch_size
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
reshape_q = q.view(bs, -1, layer.tp_q_head_num, layer.head_dim) reshape_q = q.view(bs, -1, layer.tp_q_head_num, layer.head_dim)
if self.is_fp8_kvcache: if self.is_fp8_kvcache:
@@ -489,10 +489,10 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
if k is not None: if k is not None:
assert v is not None assert v is not None
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
bs = forward_batch.batch_size bs = forward_batch.batch_size
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
reshape_q = q.view(bs, -1, layer.tp_q_head_num, layer.head_dim) reshape_q = q.view(bs, -1, layer.tp_q_head_num, layer.head_dim)
if self.is_fp8_kvcache: if self.is_fp8_kvcache:
@@ -23,6 +23,8 @@ class HybridAttnBackend(AttentionBackend):
self.prefill_backend = prefill_backend self.prefill_backend = prefill_backend
self.decode_backend = decode_backend self.decode_backend = decode_backend
self.data_type = model_runner.kv_cache_dtype self.data_type = model_runner.kv_cache_dtype
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.req_to_token_pool = model_runner.req_to_token_pool
def _select_backend(self, forward_mode: ForwardMode) -> AttentionBackend: def _select_backend(self, forward_mode: ForwardMode) -> AttentionBackend:
""" """
@@ -143,6 +143,7 @@ class MambaAttnBackendBase(AttentionBackend):
self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.topk = model_runner.server_args.speculative_eagle_topk or 0
self.is_draft_worker = model_runner.is_draft_worker self.is_draft_worker = model_runner.is_draft_worker
self.req_to_token_pool: HybridReqToTokenPool = model_runner.req_to_token_pool self.req_to_token_pool: HybridReqToTokenPool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.forward_metadata: ForwardMetadata = None self.forward_metadata: ForwardMetadata = None
self.state_indices_list = [] self.state_indices_list = []
self.query_start_loc_list = [] self.query_start_loc_list = []
@@ -763,6 +764,9 @@ class HybridLinearAttnBackend(AttentionBackend):
self.full_attn_backend = full_attn_backend self.full_attn_backend = full_attn_backend
self.linear_attn_backend = linear_attn_backend self.linear_attn_backend = linear_attn_backend
self.attn_backend_list = [full_attn_backend, linear_attn_backend] self.attn_backend_list = [full_attn_backend, linear_attn_backend]
# Dispatcher aliases the full-attn backend's pool refs.
self.token_to_kv_pool = full_attn_backend.token_to_kv_pool
self.req_to_token_pool = full_attn_backend.req_to_token_pool
def _is_full_attn( def _is_full_attn(
self, layer: Optional[RadixAttention], layer_id: Optional[int] = None self, layer: Optional[RadixAttention], layer_id: Optional[int] = None
@@ -19,6 +19,10 @@ class IntelAMXAttnBackend(AttentionBackend):
super().__init__() super().__init__()
self.forward_metadata = None self.forward_metadata = None
self.device = model_runner.device self.device = model_runner.device
# Pool refs — captured at construction so they survive deletion of the
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.num_head = ( self.num_head = (
model_runner.model_config.num_attention_heads // model_runner.tp_size model_runner.model_config.num_attention_heads // model_runner.tp_size
@@ -105,7 +109,7 @@ class IntelAMXAttnBackend(AttentionBackend):
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
if save_kv_cache and k is not None and v is not None: if save_kv_cache and k is not None and v is not None:
forward_batch.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
_, max_extend_len = self.forward_metadata _, max_extend_len = self.forward_metadata
self.extend_attention_fwd( self.extend_attention_fwd(
@@ -113,9 +117,9 @@ class IntelAMXAttnBackend(AttentionBackend):
k, k,
v, v,
o.view(-1, layer.tp_q_head_num, layer.v_head_dim), o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), self.token_to_kv_pool.get_key_buffer(layer.layer_id),
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id), self.token_to_kv_pool.get_value_buffer(layer.layer_id),
forward_batch.req_to_token_pool.req_to_token, self.req_to_token_pool.req_to_token,
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.seq_lens, forward_batch.seq_lens,
forward_batch.extend_seq_lens, forward_batch.extend_seq_lens,
@@ -152,14 +156,14 @@ class IntelAMXAttnBackend(AttentionBackend):
) )
self.decode_attention_fwd( self.decode_attention_fwd(
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), self.token_to_kv_pool.get_key_buffer(layer.layer_id),
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id), self.token_to_kv_pool.get_value_buffer(layer.layer_id),
o.view(-1, layer.tp_q_head_num, layer.v_head_dim), o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
k, k,
v, v,
cache_loc, cache_loc,
attn_logits, attn_logits,
forward_batch.req_to_token_pool.req_to_token, self.req_to_token_pool.req_to_token,
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.seq_lens, forward_batch.seq_lens,
layer.scaling, layer.scaling,
@@ -15,6 +15,10 @@ class TboAttnBackend(AttentionBackend):
super().__init__() super().__init__()
self.primary = primary self.primary = primary
self.children = children self.children = children
# Dispatcher aliases the primary's pool refs so get_attn_backend()
# reads through TboAttnBackend resolve to the underlying pool.
self.token_to_kv_pool = primary.token_to_kv_pool
self.req_to_token_pool = primary.req_to_token_pool
@classmethod @classmethod
def init_new(cls, creator: Callable[[], AttentionBackend]): def init_new(cls, creator: Callable[[], AttentionBackend]):
@@ -249,7 +249,7 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
# reproduces the original [tokens, 1, qk_rope] latent layout. # reproduces the original [tokens, 1, qk_rope] latent layout.
kv_a_fp8 = fp8_quantize(kv_a, enable_pdl=is_arch_support_pdl()) kv_a_fp8 = fp8_quantize(kv_a, enable_pdl=is_arch_support_pdl())
k_pe_fp8 = k_fp8[:, 0:1, layer.qk_nope_head_dim :] k_pe_fp8 = k_fp8[:, 0:1, layer.qk_nope_head_dim :]
forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.token_to_kv_pool.set_mla_kv_buffer(
layer.attn_mha, layer.attn_mha,
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
kv_a_fp8.unsqueeze(1), kv_a_fp8.unsqueeze(1),
@@ -19,6 +19,10 @@ class TorchFlexAttnBackend(AttentionBackend):
super().__init__() super().__init__()
self.forward_metadata = None self.forward_metadata = None
self.device = model_runner.device self.device = model_runner.device
# Pool refs — captured at construction so they survive deletion of the
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.flex_attention = torch.compile(flex_attention, dynamic=True) self.flex_attention = torch.compile(flex_attention, dynamic=True)
torch._dynamo.config.cache_size_limit = 1024 torch._dynamo.config.cache_size_limit = 1024
torch._dynamo.config.accumulated_cache_size_limit = 1024 torch._dynamo.config.accumulated_cache_size_limit = 1024
@@ -248,7 +252,7 @@ class TorchFlexAttnBackend(AttentionBackend):
o = torch.empty_like(q) o = torch.empty_like(q)
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v layer, forward_batch.out_cache_loc, k, v
) )
@@ -266,9 +270,9 @@ class TorchFlexAttnBackend(AttentionBackend):
self._run_flex_forward_extend( self._run_flex_forward_extend(
q_, q_,
o_, o_,
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), self.token_to_kv_pool.get_key_buffer(layer.layer_id),
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id), self.token_to_kv_pool.get_value_buffer(layer.layer_id),
forward_batch.req_to_token_pool.req_to_token, self.req_to_token_pool.req_to_token,
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.seq_lens, forward_batch.seq_lens,
forward_batch.extend_prefix_lens, forward_batch.extend_prefix_lens,
@@ -298,7 +302,7 @@ class TorchFlexAttnBackend(AttentionBackend):
o = torch.empty_like(q) o = torch.empty_like(q)
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v layer, forward_batch.out_cache_loc, k, v
) )
@@ -309,9 +313,9 @@ class TorchFlexAttnBackend(AttentionBackend):
self._run_flex_forward_decode( self._run_flex_forward_decode(
q_, q_,
o_, o_,
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), self.token_to_kv_pool.get_key_buffer(layer.layer_id),
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id), self.token_to_kv_pool.get_value_buffer(layer.layer_id),
forward_batch.req_to_token_pool.req_to_token, self.req_to_token_pool.req_to_token,
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.seq_lens, forward_batch.seq_lens,
scaling=layer.scaling, scaling=layer.scaling,
@@ -19,6 +19,10 @@ class TorchNativeAttnBackend(AttentionBackend):
super().__init__() super().__init__()
self.forward_metadata = None self.forward_metadata = None
self.device = model_runner.device self.device = model_runner.device
# Pool refs — captured at construction so they survive deletion of the
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init the metadata for a forward pass.""" """Init the metadata for a forward pass."""
@@ -235,7 +239,7 @@ class TorchNativeAttnBackend(AttentionBackend):
cache_loc = forward_batch.out_cache_loc cache_loc = forward_batch.out_cache_loc
if save_kv_cache and k is not None and v is not None: if save_kv_cache and k is not None and v is not None:
forward_batch.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
use_gqa = layer.tp_q_head_num != layer.tp_k_head_num use_gqa = layer.tp_q_head_num != layer.tp_k_head_num
@@ -249,9 +253,9 @@ class TorchNativeAttnBackend(AttentionBackend):
self._run_sdpa_forward_extend( self._run_sdpa_forward_extend(
q_, q_,
o_, o_,
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), self.token_to_kv_pool.get_key_buffer(layer.layer_id),
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id), self.token_to_kv_pool.get_value_buffer(layer.layer_id),
forward_batch.req_to_token_pool.req_to_token, self.req_to_token_pool.req_to_token,
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.seq_lens, forward_batch.seq_lens,
forward_batch.extend_prefix_lens, forward_batch.extend_prefix_lens,
@@ -292,9 +296,8 @@ class TorchNativeAttnBackend(AttentionBackend):
else: else:
cache_loc = forward_batch.out_cache_loc cache_loc = forward_batch.out_cache_loc
if save_kv_cache: if save_kv_cache and k is not None and v is not None:
if k is not None and v is not None: self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
forward_batch.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
use_gqa = layer.tp_q_head_num != layer.tp_k_head_num use_gqa = layer.tp_q_head_num != layer.tp_k_head_num
@@ -304,9 +307,9 @@ class TorchNativeAttnBackend(AttentionBackend):
self._run_sdpa_forward_decode( self._run_sdpa_forward_decode(
q_, q_,
o_, o_,
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), self.token_to_kv_pool.get_key_buffer(layer.layer_id),
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id), self.token_to_kv_pool.get_value_buffer(layer.layer_id),
forward_batch.req_to_token_pool.req_to_token, self.req_to_token_pool.req_to_token,
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.seq_lens, forward_batch.seq_lens,
forward_batch.encoder_lens, forward_batch.encoder_lens,
@@ -105,6 +105,10 @@ class TritonAttnBackend(AttentionBackend):
self.skip_prefill = skip_prefill self.skip_prefill = skip_prefill
max_bs = model_runner.req_to_token_pool.size max_bs = model_runner.req_to_token_pool.size
self.sliding_window_size = model_runner.sliding_window_size self.sliding_window_size = model_runner.sliding_window_size
# Pool refs — captured at construction so they survive deletion of the
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.req_to_token = model_runner.req_to_token_pool.req_to_token self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.token_to_kv_pool_allocator = model_runner.token_to_kv_pool_allocator self.token_to_kv_pool_allocator = model_runner.token_to_kv_pool_allocator
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
@@ -904,7 +908,7 @@ class TritonAttnBackend(AttentionBackend):
o = torch.empty_like(q) o = torch.empty_like(q)
if k is None and v is None: if k is None and v is None:
pool = forward_batch.token_to_kv_pool pool = self.token_to_kv_pool
cache_loc = forward_batch.out_cache_loc cache_loc = forward_batch.out_cache_loc
if isinstance(pool, SWAKVPool) and pool.layers_mapping[layer.layer_id][1]: if isinstance(pool, SWAKVPool) and pool.layers_mapping[layer.layer_id][1]:
cache_loc = pool.translate_loc_from_full_to_swa(cache_loc) cache_loc = pool.translate_loc_from_full_to_swa(cache_loc)
@@ -917,7 +921,7 @@ class TritonAttnBackend(AttentionBackend):
# Save KV cache first (must do this before unified kernel) # Save KV cache first (must do this before unified kernel)
if save_kv_cache: if save_kv_cache:
if layer.k_scale is None: if layer.k_scale is None:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, layer,
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
k, k,
@@ -928,14 +932,14 @@ class TritonAttnBackend(AttentionBackend):
# doesn't accept scale parameters. Clone to protect k from mutation # doesn't accept scale parameters. Clone to protect k from mutation
# since it's used later in the attention kernel. # since it's used later in the attention kernel.
k_scaled = k.clone().div_(layer.k_scale) k_scaled = k.clone().div_(layer.k_scale)
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, layer,
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
k_scaled, k_scaled,
v, v,
) )
else: else:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, layer,
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
k.clone(), # cloned to protect k,v from in-place mutation in set_kv_buffer k.clone(), # cloned to protect k,v from in-place mutation in set_kv_buffer
@@ -989,8 +993,8 @@ class TritonAttnBackend(AttentionBackend):
k.contiguous(), k.contiguous(),
v.contiguous(), v.contiguous(),
o.view(-1, layer.tp_q_head_num, layer.v_head_dim), o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), self.token_to_kv_pool.get_key_buffer(layer.layer_id),
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id), self.token_to_kv_pool.get_value_buffer(layer.layer_id),
self.forward_metadata.qo_indptr, self.forward_metadata.qo_indptr,
kv_indptr, kv_indptr,
kv_indices, kv_indices,
@@ -1058,7 +1062,7 @@ class TritonAttnBackend(AttentionBackend):
window_start_pos = None window_start_pos = None
extend_kv_indices = forward_batch.out_cache_loc extend_kv_indices = forward_batch.out_cache_loc
pool = forward_batch.token_to_kv_pool pool = self.token_to_kv_pool
if ( if (
layer.sliding_window_size is not None layer.sliding_window_size is not None
and layer.sliding_window_size > -1 and layer.sliding_window_size > -1
@@ -1124,8 +1128,8 @@ class TritonAttnBackend(AttentionBackend):
self.extend_attention_fwd_unified( self.extend_attention_fwd_unified(
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
o.view(-1, layer.tp_q_head_num, layer.v_head_dim), o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), self.token_to_kv_pool.get_key_buffer(layer.layer_id),
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id), self.token_to_kv_pool.get_value_buffer(layer.layer_id),
k_descale, k_descale,
v_descale, v_descale,
self.forward_metadata.qo_indptr, self.forward_metadata.qo_indptr,
@@ -1174,14 +1178,14 @@ class TritonAttnBackend(AttentionBackend):
# MLATokenToKVPool doesn't accept scale parameters; k is unused # MLATokenToKVPool doesn't accept scale parameters; k is unused
# after this point in decode, so scale in place. # after this point in decode, so scale in place.
k.div_(layer.k_scale) k.div_(layer.k_scale)
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, layer,
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
k, k,
v, v,
) )
else: else:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, layer,
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
k, k,
@@ -1216,8 +1220,8 @@ class TritonAttnBackend(AttentionBackend):
self.decode_attention_fwd( self.decode_attention_fwd(
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), self.token_to_kv_pool.get_key_buffer(layer.layer_id),
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id), self.token_to_kv_pool.get_value_buffer(layer.layer_id),
o.view(-1, layer.tp_q_head_num, layer.v_head_dim), o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
kv_indptr, kv_indptr,
kv_indices, kv_indices,
@@ -1275,6 +1279,7 @@ class TritonMultiStepDraftBackend:
) )
self.device = model_runner.device self.device = model_runner.device
# Cached variables for generate_draft_decode_kv_indices # Cached variables for generate_draft_decode_kv_indices
self.req_to_token_pool = model_runner.req_to_token_pool
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1] self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
self.page_size = model_runner.server_args.page_size self.page_size = model_runner.server_args.page_size
@@ -1295,7 +1300,7 @@ class TritonMultiStepDraftBackend:
(self.speculative_num_steps, num_seqs, self.topk) (self.speculative_num_steps, num_seqs, self.topk)
]( ](
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.req_to_token_pool.req_to_token, self.req_to_token_pool.req_to_token,
forward_batch.seq_lens, forward_batch.seq_lens,
kv_indices_buffer, kv_indices_buffer,
self.kv_indptr, self.kv_indptr,
@@ -558,7 +558,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
cache_loc = self._get_layer_cache_loc(layer, forward_batch) cache_loc = self._get_layer_cache_loc(layer, forward_batch)
# Get K/V cache buffers from token_to_kv_pool # Get K/V cache buffers from token_to_kv_pool
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id) k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
fused_fp8_set_kv_buffer( fused_fp8_set_kv_buffer(
k=k, k=k,
@@ -598,7 +598,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
), ),
(1, 0), (1, 0),
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
else: else:
@@ -611,7 +611,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
metadata.cu_seqlens_k = torch.nn.functional.pad( metadata.cu_seqlens_k = torch.nn.functional.pad(
torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0) torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0)
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
elif forward_batch.forward_mode.is_target_verify(): elif forward_batch.forward_mode.is_target_verify():
@@ -635,7 +635,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32), torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32),
(1, 0), (1, 0),
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
@@ -645,7 +645,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
metadata.cu_seqlens_k = torch.nn.functional.pad( metadata.cu_seqlens_k = torch.nn.functional.pad(
torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0) torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0)
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
@@ -713,7 +713,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
else: else:
# Use original set_kv_buffer path # Use original set_kv_buffer path
if save_kv_cache and k is not None: if save_kv_cache and k is not None:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
@@ -721,7 +721,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
if self.data_type == torch.float8_e4m3fn and (not self.is_xqa_impl): if self.data_type == torch.float8_e4m3fn and (not self.is_xqa_impl):
q = q.to(torch.float8_e4m3fn) q = q.to(torch.float8_e4m3fn)
q = q.reshape(-1, layer.tp_q_head_num, layer.head_dim) q = q.reshape(-1, layer.tp_q_head_num, layer.head_dim)
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id) k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
# shape conversion: # shape conversion:
# [num_pages, page_size, num_kv_heads, head_dim] -> [num_pages, num_kv_heads, page_size, head_dim] # [num_pages, page_size, num_kv_heads, head_dim] -> [num_pages, num_kv_heads, page_size, head_dim]
k_cache = k_cache.view( k_cache = k_cache.view(
@@ -799,7 +799,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
else: else:
# Use original set_kv_buffer path # Use original set_kv_buffer path
if save_kv_cache and k is not None: if save_kv_cache and k is not None:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
@@ -807,7 +807,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
q = q.to(torch.float8_e4m3fn) q = q.to(torch.float8_e4m3fn)
q = q.reshape(-1, layer.tp_q_head_num, layer.head_dim) q = q.reshape(-1, layer.tp_q_head_num, layer.head_dim)
# [num_pages, page_size, num_kv_heads, head_dim] -> [num_pages, num_kv_heads, page_size, head_dim] # [num_pages, page_size, num_kv_heads, head_dim] -> [num_pages, num_kv_heads, page_size, head_dim]
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id) k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
k_cache = k_cache.view( k_cache = k_cache.view(
-1, self.page_size, layer.tp_k_head_num, layer.head_dim -1, self.page_size, layer.tp_k_head_num, layer.head_dim
).permute(0, 2, 1, 3) ).permute(0, 2, 1, 3)
@@ -898,7 +898,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
assert ( assert (
k is not None and k_rope is not None k is not None and k_rope is not None
), "For populating trtllm_mla kv cache, both k_nope and k_rope should be not None." ), "For populating trtllm_mla kv cache, both k_nope and k_rope should be not None."
forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.token_to_kv_pool.set_mla_kv_buffer(
layer, forward_batch.out_cache_loc, k, k_rope layer, forward_batch.out_cache_loc, k, k_rope
) )
@@ -924,7 +924,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
query = query.unsqueeze(1) query = query.unsqueeze(1)
# Prepare KV cache inline # Prepare KV cache inline
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
kv_cache = k_cache.view(-1, self.page_size, self.kv_cache_dim).unsqueeze(1) kv_cache = k_cache.view(-1, self.page_size, self.kv_cache_dim).unsqueeze(1)
# Get metadata # Get metadata
@@ -1005,7 +1005,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
assert ( assert (
k is not None and k_rope is not None k is not None and k_rope is not None
), "For populating trtllm_mla kv cache, both k_nope and k_rope should be not None." ), "For populating trtllm_mla kv cache, both k_nope and k_rope should be not None."
forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.token_to_kv_pool.set_mla_kv_buffer(
layer, forward_batch.out_cache_loc, k, k_rope layer, forward_batch.out_cache_loc, k, k_rope
) )
@@ -1046,7 +1046,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
# Ensure query has shape [bs, num_draft_tokens, num_q_heads, head_dim] # Ensure query has shape [bs, num_draft_tokens, num_q_heads, head_dim]
bs = forward_batch.batch_size bs = forward_batch.batch_size
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
kv_cache = k_cache.view(-1, self.page_size, self.kv_cache_dim).unsqueeze(1) kv_cache = k_cache.view(-1, self.page_size, self.kv_cache_dim).unsqueeze(1)
q = q.to(self.data_type) q = q.to(self.data_type)
@@ -117,6 +117,11 @@ class WaveAttnBackend(AttentionBackend):
self.skip_prefill = skip_prefill self.skip_prefill = skip_prefill
# Pool refs — captured at construction so they survive deletion of the
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
max_bs = model_runner.req_to_token_pool.size max_bs = model_runner.req_to_token_pool.size
if kv_indptr_buf is None: if kv_indptr_buf is None:
@@ -556,7 +561,7 @@ class WaveAttnBackend(AttentionBackend):
o = torch.empty_like(q) o = torch.empty_like(q)
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v layer, forward_batch.out_cache_loc, k, v
) )
@@ -571,8 +576,8 @@ class WaveAttnBackend(AttentionBackend):
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
k.contiguous(), k.contiguous(),
v.contiguous(), v.contiguous(),
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), self.token_to_kv_pool.get_key_buffer(layer.layer_id),
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id), self.token_to_kv_pool.get_value_buffer(layer.layer_id),
self.forward_metadata.qo_indptr, self.forward_metadata.qo_indptr,
self.forward_metadata.kv_indptr, self.forward_metadata.kv_indptr,
self.forward_metadata.kv_indices, self.forward_metadata.kv_indices,
@@ -606,14 +611,14 @@ class WaveAttnBackend(AttentionBackend):
o = torch.empty_like(q) o = torch.empty_like(q)
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, forward_batch.out_cache_loc, k, v layer, forward_batch.out_cache_loc, k, v
) )
self.decode_attention_fwd( self.decode_attention_fwd(
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), self.token_to_kv_pool.get_key_buffer(layer.layer_id),
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id), self.token_to_kv_pool.get_value_buffer(layer.layer_id),
o.view(-1, layer.tp_q_head_num, layer.v_head_dim), o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
self.forward_metadata.kv_indptr, self.forward_metadata.kv_indptr,
self.forward_metadata.kv_indices, self.forward_metadata.kv_indices,
@@ -61,6 +61,10 @@ class XPUAttentionBackend(AttentionBackend):
self.device = model_runner.device self.device = model_runner.device
self.decode_cuda_graph_metadata = {} self.decode_cuda_graph_metadata = {}
self.target_verify_metadata = {} self.target_verify_metadata = {}
# Pool refs — captured at construction so they survive deletion of the
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.req_to_token = model_runner.req_to_token_pool.req_to_token self.req_to_token = model_runner.req_to_token_pool.req_to_token
self.kv_cache_dtype = model_runner.kv_cache_dtype self.kv_cache_dtype = model_runner.kv_cache_dtype
self.kv_cache_dtype_str = model_runner.server_args.kv_cache_dtype self.kv_cache_dtype_str = model_runner.server_args.kv_cache_dtype
@@ -122,7 +126,7 @@ class XPUAttentionBackend(AttentionBackend):
), ),
(1, 0), (1, 0),
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
else: else:
@@ -142,7 +146,7 @@ class XPUAttentionBackend(AttentionBackend):
), ),
(1, 0), (1, 0),
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
@@ -186,7 +190,7 @@ class XPUAttentionBackend(AttentionBackend):
metadata.cu_seqlens_k = torch.nn.functional.pad( metadata.cu_seqlens_k = torch.nn.functional.pad(
torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0) torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0)
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
# TODO: we need to test this part for llama 4 eagle case # TODO: we need to test this part for llama 4 eagle case
@@ -214,7 +218,7 @@ class XPUAttentionBackend(AttentionBackend):
), ),
(1, 0), (1, 0),
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
@@ -236,7 +240,7 @@ class XPUAttentionBackend(AttentionBackend):
), ),
(1, 0), (1, 0),
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
@@ -297,7 +301,7 @@ class XPUAttentionBackend(AttentionBackend):
) )
_, sort_order = torch.sort(keys, dim=1) _, sort_order = torch.sort(keys, dim=1)
non_masked_page_table = ( non_masked_page_table = (
forward_batch.req_to_token_pool.req_to_token[ self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : forward_batch.req_pool_indices, :
] ]
.gather(1, cols) .gather(1, cols)
@@ -324,7 +328,7 @@ class XPUAttentionBackend(AttentionBackend):
metadata.cu_seqlens_k = torch.nn.functional.pad( metadata.cu_seqlens_k = torch.nn.functional.pad(
torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0) torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0)
) )
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.max_seq_len_k forward_batch.req_pool_indices, : metadata.max_seq_len_k
] ]
@@ -357,12 +361,12 @@ class XPUAttentionBackend(AttentionBackend):
(1, 0), (1, 0),
) )
metadata.encoder_max_seq_len_k = metadata.encoder_lens_int32.max().item() metadata.encoder_max_seq_len_k = metadata.encoder_lens_int32.max().item()
metadata.encoder_page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.encoder_page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, : metadata.encoder_max_seq_len_k forward_batch.req_pool_indices, : metadata.encoder_max_seq_len_k
] ]
# Currently only support forward_batch.encoder_lens.numel() == 1 # Currently only support forward_batch.encoder_lens.numel() == 1
metadata.page_table = forward_batch.req_to_token_pool.req_to_token[ metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
metadata.encoder_max_seq_len_k : ( metadata.encoder_max_seq_len_k : (
metadata.encoder_max_seq_len_k + metadata.max_seq_len_k metadata.encoder_max_seq_len_k + metadata.max_seq_len_k
@@ -418,11 +422,11 @@ class XPUAttentionBackend(AttentionBackend):
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
if not self.use_mla: if not self.use_mla:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
else: else:
forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.token_to_kv_pool.set_mla_kv_buffer(
layer, layer,
cache_loc, cache_loc,
k, k,
@@ -501,9 +505,7 @@ class XPUAttentionBackend(AttentionBackend):
# Use Flash Attention for prefill # Use Flash Attention for prefill
if not self.use_mla: if not self.use_mla:
# Do multi-head attention # Do multi-head attention
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer( key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
layer.layer_id
)
key_cache = key_cache.view( key_cache = key_cache.view(
-1, self.page_size, layer.tp_k_head_num, layer.head_dim -1, self.page_size, layer.tp_k_head_num, layer.head_dim
) )
@@ -614,9 +616,9 @@ class XPUAttentionBackend(AttentionBackend):
return output return output
else: else:
# Do absorbed multi-latent attention # Do absorbed multi-latent attention
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer( kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
layer.layer_id q.dtype
).to(q.dtype) )
k_rope = kv_cache[:, :, layer.v_head_dim :] k_rope = kv_cache[:, :, layer.v_head_dim :]
c_kv = kv_cache[:, :, : layer.v_head_dim] c_kv = kv_cache[:, :, : layer.v_head_dim]
k_rope_cache = k_rope.view( k_rope_cache = k_rope.view(
@@ -710,14 +712,14 @@ class XPUAttentionBackend(AttentionBackend):
else forward_batch.encoder_out_cache_loc else forward_batch.encoder_out_cache_loc
) )
if not self.use_mla: if not self.use_mla:
forward_batch.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
layer, cache_loc, k, v, layer.k_scale, layer.v_scale layer, cache_loc, k, v, layer.k_scale, layer.v_scale
) )
else: else:
k_rope_val = ( k_rope_val = (
k_rope if k_rope is not None else k[:, :, layer.v_head_dim :] k_rope if k_rope is not None else k[:, :, layer.v_head_dim :]
) )
forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.token_to_kv_pool.set_mla_kv_buffer(
layer, layer,
cache_loc, cache_loc,
k, k,
@@ -768,9 +770,7 @@ class XPUAttentionBackend(AttentionBackend):
if not self.use_mla: if not self.use_mla:
# Do multi-head attention # Do multi-head attention
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer( key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
layer.layer_id
)
key_cache = key_cache.view( key_cache = key_cache.view(
-1, self.page_size, layer.tp_k_head_num, layer.head_dim -1, self.page_size, layer.tp_k_head_num, layer.head_dim
) )
@@ -876,9 +876,7 @@ class XPUAttentionBackend(AttentionBackend):
o = result o = result
else: else:
# Do absorbed multi-latent attention # Do absorbed multi-latent attention
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id).to( kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(q.dtype)
q.dtype
)
assert not use_cascade_attn, "Cascade attention is not supported with MLA" assert not use_cascade_attn, "Cascade attention is not supported with MLA"
if q_rope is not None: if q_rope is not None:
+3 -2
View File
@@ -29,6 +29,7 @@ from sglang.srt.model_executor.breakable_cuda_graph.breakable_cuda_graph import
from sglang.srt.model_executor.breakable_cuda_graph.context import ( from sglang.srt.model_executor.breakable_cuda_graph.context import (
is_in_breakable_cuda_graph, is_in_breakable_cuda_graph,
) )
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.custom_op import register_custom_op
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -135,7 +136,7 @@ class RadixAttention(nn.Module):
) )
return output return output
else: else:
return forward_batch.attn_backend.forward( return get_attn_backend().forward(
q, q,
k, k,
v, v,
@@ -188,7 +189,7 @@ def unified_attention_with_output(
# the FA kernel validates out.size(0) == q.size(0). # the FA kernel validates out.size(0) == q.size(0).
forward_batch._attn_output = output[:real_num_tokens] forward_batch._attn_output = output[:real_num_tokens]
ret = forward_batch.attn_backend.forward( ret = get_attn_backend().forward(
query, query,
key, key,
value, value,
@@ -22,6 +22,7 @@ from torch import nn
from sglang.srt.compilation.compilation_config import register_split_op from sglang.srt.compilation.compilation_config import register_split_op
from sglang.srt.compilation.piecewise_context_manager import get_forward_context from sglang.srt.compilation.piecewise_context_manager import get_forward_context
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.custom_op import register_custom_op
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -92,7 +93,7 @@ class RadixLinearAttention(nn.Module):
) )
return output return output
else: else:
return forward_batch.attn_backend.forward( return get_attn_backend().forward(
layer=self, layer=self,
forward_batch=forward_batch, forward_batch=forward_batch,
mixed_qkv=mixed_qkv, mixed_qkv=mixed_qkv,
@@ -124,7 +125,7 @@ def unified_linear_attention_with_output(
# this backend call so model/backend state is still written to the same batch. # this backend call so model/backend state is still written to the same batch.
forward_batch.out_cache_loc = original_out_cache_loc[:real_num_tokens] forward_batch.out_cache_loc = original_out_cache_loc[:real_num_tokens]
ret = forward_batch.attn_backend.forward( ret = get_attn_backend().forward(
layer=attention_layer, layer=attention_layer,
forward_batch=forward_batch, forward_batch=forward_batch,
mixed_qkv=mixed_qkv[:real_num_tokens], mixed_qkv=mixed_qkv[:real_num_tokens],
+2 -1
View File
@@ -14,6 +14,7 @@ from sglang.srt.layers.dp_attention import (
get_attention_cp_size, get_attention_cp_size,
is_allocation_symmetric, is_allocation_symmetric,
) )
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
@@ -342,7 +343,7 @@ def cp_allgather_and_save_kv_cache(forward_batch, layer, k, v, cp_size):
v, cp_size, forward_batch, torch.cuda.current_stream() v, cp_size, forward_batch, torch.cuda.current_stream()
) )
forward_batch.token_to_kv_pool.set_kv_buffer( get_token_to_kv_pool().set_kv_buffer(
layer, layer,
cache_loc, cache_loc,
key_cache_full, key_cache_full,
@@ -57,6 +57,7 @@ from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode, CaptureHiddenMode,
PPProxyTensors, PPProxyTensors,
) )
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.piecewise_cuda_graph_runner import ( from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
PiecewiseCudaGraphRunner, PiecewiseCudaGraphRunner,
freeze_gc, freeze_gc,
@@ -292,9 +293,6 @@ class BreakableCudaGraphRunner:
next_token_logits_buffer=None, next_token_logits_buffer=None,
orig_seq_lens=orig_seq_lens, orig_seq_lens=orig_seq_lens,
seq_lens_cpu=torch.tensor([num_tokens], device="cpu"), seq_lens_cpu=torch.tensor([num_tokens], device="cpu"),
req_to_token_pool=self.model_runner.req_to_token_pool,
token_to_kv_pool=self.model_runner.token_to_kv_pool,
attn_backend=self.model_runner.attn_backend,
out_cache_loc=buffers.out_cache_loc[:num_tokens], out_cache_loc=buffers.out_cache_loc[:num_tokens],
seq_lens_sum=num_tokens, seq_lens_sum=num_tokens,
mamba_track_indices=None, mamba_track_indices=None,
@@ -329,8 +327,11 @@ class BreakableCudaGraphRunner:
"""Warmup the model with a forward pass.""" """Warmup the model with a forward pass."""
num_tokens = self.capture_num_tokens[0] num_tokens = self.capture_num_tokens[0]
forward_batch = self._build_capture_forward_batch(num_tokens) forward_batch = self._build_capture_forward_batch(num_tokens)
self.model_runner.attn_backend.init_forward_metadata(forward_batch) with forward_context(
self._run_forward(forward_batch, num_tokens) ForwardContext(attn_backend=self.model_runner.attn_backend)
):
self.model_runner.attn_backend.init_forward_metadata(forward_batch)
self._run_forward(forward_batch, num_tokens)
def _capture_all(self): def _capture_all(self):
"""Capture breakable CUDA graphs for all token sizes.""" """Capture breakable CUDA graphs for all token sizes."""
@@ -394,14 +395,17 @@ class BreakableCudaGraphRunner:
self.model_runner.token_to_kv_pool.invalidate_loc_cache() self.model_runner.token_to_kv_pool.invalidate_loc_cache()
return self._run_forward(forward_batch, num_tokens) return self._run_forward(forward_batch, num_tokens)
for _ in range(2): with forward_context(
self.device_module.synchronize() ForwardContext(attn_backend=self.model_runner.attn_backend)
self.model_runner.tp_group.barrier() ):
run_once() for _ in range(2):
self.device_module.synchronize()
self.model_runner.tp_group.barrier()
run_once()
graph = BreakableCUDAGraph() graph = BreakableCUDAGraph()
with BreakableCUDAGraphCapture(cuda_graph=graph, pool=pool, stream=stream): with BreakableCUDAGraphCapture(cuda_graph=graph, pool=pool, stream=stream):
output = run_once() output = run_once()
return graph, output return graph, output
@@ -36,6 +36,7 @@ from sglang.srt.model_executor.forward_batch_info import (
PPProxyTensors, PPProxyTensors,
enable_num_token_non_padded, enable_num_token_non_padded,
) )
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.utils import ( from sglang.srt.utils import (
log_info_on_rank0, log_info_on_rank0,
require_attn_tp_gather, require_attn_tp_gather,
@@ -679,9 +680,6 @@ class CPUGraphRunner:
input_ids=input_ids, input_ids=input_ids,
req_pool_indices=req_pool_indices, req_pool_indices=req_pool_indices,
seq_lens=seq_lens, seq_lens=seq_lens,
req_to_token_pool=self.model_runner.req_to_token_pool,
token_to_kv_pool=self.model_runner.token_to_kv_pool,
attn_backend=self.model_runner.attn_backend,
out_cache_loc=out_cache_loc, out_cache_loc=out_cache_loc,
seq_lens_sum=seq_lens.sum().item(), seq_lens_sum=seq_lens.sum().item(),
return_logprob=False, return_logprob=False,
@@ -693,43 +691,46 @@ class CPUGraphRunner:
num_token_non_padded=self.num_token_non_padded, num_token_non_padded=self.num_token_non_padded,
global_forward_mode=self.capture_forward_mode, global_forward_mode=self.capture_forward_mode,
) )
self.model_runner.attn_backend.init_forward_metadata_capture_cpu_graph( with forward_context(
bs, ForwardContext(attn_backend=self.model_runner.attn_backend)
num_tokens, ):
req_pool_indices, self.model_runner.attn_backend.init_forward_metadata_capture_cpu_graph(
seq_lens, bs,
None, num_tokens,
forward_batch.forward_mode, req_pool_indices,
forward_batch.spec_info, seq_lens,
) None,
# Do infernence to avoid setting attr at runtime, e.g., forward_batch.forward_mode,
# self.attn_mha.kv_b_proj = self.kv_b_proj for full graph compile on CPU forward_batch.spec_info,
with torch.no_grad():
self.model_runner.tp_group.barrier()
self.model_runner.model.forward(
forward_batch.input_ids,
forward_batch.positions,
forward_batch,
) )
with torch.no_grad():
# 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
logits_output_or_pp_proxy_tensors = forward(
forward_batch.input_ids,
forward_batch.positions,
forward_batch,
)
return logits_output_or_pp_proxy_tensors
with torch.no_grad():
for _ in range(2):
self.model_runner.tp_group.barrier() self.model_runner.tp_group.barrier()
out = run_once() self.model_runner.model.forward(
# Save the captured forward_batch forward_batch.input_ids,
self.captured_forward_batches[bs] = forward_batch forward_batch.positions,
return forward, out forward_batch,
)
# 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
)
logits_output_or_pp_proxy_tensors = forward(
forward_batch.input_ids,
forward_batch.positions,
forward_batch,
)
return logits_output_or_pp_proxy_tensors
with torch.no_grad():
for _ in range(2):
self.model_runner.tp_group.barrier()
out = run_once()
# Save the captured forward_batch
self.captured_forward_batches[bs] = forward_batch
return forward, out
def recapture_if_needed(self, forward_batch: ForwardBatch): def recapture_if_needed(self, forward_batch: ForwardBatch):
@@ -65,6 +65,7 @@ from sglang.srt.model_executor.forward_batch_info import (
compute_local_num_token_non_padded, compute_local_num_token_non_padded,
enable_num_token_non_padded, enable_num_token_non_padded,
) )
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups
from sglang.srt.utils import ( from sglang.srt.utils import (
@@ -1016,9 +1017,6 @@ class CudaGraphRunner:
seq_lens_cpu=seq_lens_cpu, seq_lens_cpu=seq_lens_cpu,
next_token_logits_buffer=next_token_logits_buffer, next_token_logits_buffer=next_token_logits_buffer,
orig_seq_lens=seq_lens, orig_seq_lens=seq_lens,
req_to_token_pool=self.model_runner.req_to_token_pool,
token_to_kv_pool=self.model_runner.token_to_kv_pool,
attn_backend=attn_backend,
out_cache_loc=out_cache_loc, out_cache_loc=out_cache_loc,
seq_lens_sum=seq_lens.sum().item(), seq_lens_sum=seq_lens.sum().item(),
mamba_track_indices=mamba_track_indices, mamba_track_indices=mamba_track_indices,
@@ -1040,85 +1038,90 @@ class CudaGraphRunner:
lora_ids=lora_ids, lora_ids=lora_ids,
) )
# HiSparse: set coordinator so the hisparse code path is captured into the graph # Trip the coordinator so the hisparse code path is captured into the
forward_batch.hisparse_coordinator = self.model_runner.hisparse_coordinator # graph; backends read it from self.model_runner.hisparse_coordinator.
if forward_batch.hisparse_coordinator is not None: hisparse_coordinator = self.model_runner.hisparse_coordinator
forward_batch.hisparse_coordinator.num_real_reqs.fill_(bs) if hisparse_coordinator is not None:
hisparse_coordinator.num_real_reqs.fill_(bs)
if buffers.ngram_embedding_info is not None: if buffers.ngram_embedding_info is not None:
forward_batch.ngram_embedding_info = buffers.ngram_embedding_info.slice(bs) forward_batch.ngram_embedding_info = buffers.ngram_embedding_info.slice(bs)
self.tbo_plugin.capture_one_batch_size(forward_batch, num_tokens=num_tokens) # All setup hooks below read get_attn_backend() (TboForwardBatchPreparer,
# DeepEP adapter, …) so they must run inside the same ForwardContext
# that wraps the warmup/capture forward.
with forward_context(ForwardContext(attn_backend=attn_backend)):
self.tbo_plugin.capture_one_batch_size(forward_batch, num_tokens=num_tokens)
if lora_ids is not None: if lora_ids is not None:
self.model_runner.lora_manager.prepare_lora_batch(forward_batch) self.model_runner.lora_manager.prepare_lora_batch(forward_batch)
# Attention backend attn_backend.init_forward_metadata_capture_cuda_graph(
attn_backend.init_forward_metadata_capture_cuda_graph( bs,
bs,
num_tokens,
req_pool_indices,
seq_lens,
encoder_lens,
forward_batch.forward_mode,
forward_batch.spec_info,
)
# Run and capture
def run_once():
# Without this, warmup-1 caches the translation; the capture run gets
# a hit, skips the gather, and replay reuses stale SWA locations.
if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache()
# Clean intermediate result cache for DP attention
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len(
global_dp_buffer_len,
num_tokens, num_tokens,
forward_batch.dp_padding_mode.is_max_len(), req_pool_indices,
seq_lens,
encoder_lens,
forward_batch.forward_mode,
forward_batch.spec_info,
) )
set_is_extend_in_batch(False)
kwargs = {} def run_once():
if ( # Without this, warmup-1 caches the translation; the capture
self.pp_size > 1 # run hits the cache, skips the gather, and replay reuses
and "pp_proxy_tensors" in inspect.signature(forward).parameters # stale SWA locations.
): if self.model_runner.is_hybrid_swa:
kwargs["pp_proxy_tensors"] = PPProxyTensors( self.model_runner.token_to_kv_pool.invalidate_loc_cache()
{k: v.clone() for k, v in pp_proxy_tensors.tensors.items()}
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = (
None
) )
if ( set_dp_buffer_len(
self.model_runner.spec_algorithm.is_dflash() global_dp_buffer_len,
and self.model_runner.is_draft_worker num_tokens,
and "input_embeds" in inspect.signature(forward).parameters forward_batch.dp_padding_mode.is_max_len(),
): )
kwargs["input_embeds"] = buffers.input_embeds[:num_tokens] set_is_extend_in_batch(False)
logits_output_or_pp_proxy_tensors = forward( kwargs = {}
input_ids, if (
forward_batch.positions, self.pp_size > 1
forward_batch, and "pp_proxy_tensors" in inspect.signature(forward).parameters
**kwargs, ):
kwargs["pp_proxy_tensors"] = PPProxyTensors(
{k: v.clone() for k, v in pp_proxy_tensors.tensors.items()}
)
if (
self.model_runner.spec_algorithm.is_dflash()
and self.model_runner.is_draft_worker
and "input_embeds" in inspect.signature(forward).parameters
):
kwargs["input_embeds"] = buffers.input_embeds[:num_tokens]
logits_output_or_pp_proxy_tensors = forward(
input_ids,
forward_batch.positions,
forward_batch,
**kwargs,
)
return logits_output_or_pp_proxy_tensors
self.deepep_adapter.capture(is_extend_in_batch=False)
for _ in range(2):
self.device_module.synchronize()
self.model_runner.tp_group.barrier()
run_once()
attn_backend.on_after_cuda_graph_warmup()
if get_global_graph_memory_pool() is None:
set_global_graph_memory_pool(self.device_module.graph_pool_handle())
# Set graph pool id globally to be able to use symmetric memory
set_graph_pool_id(get_global_graph_memory_pool())
out = self._capture_graph(
graph, get_global_graph_memory_pool(), stream, run_once
) )
return logits_output_or_pp_proxy_tensors
self.deepep_adapter.capture(is_extend_in_batch=False)
for _ in range(2):
self.device_module.synchronize()
self.model_runner.tp_group.barrier()
run_once()
attn_backend.on_after_cuda_graph_warmup()
if get_global_graph_memory_pool() is None:
set_global_graph_memory_pool(self.device_module.graph_pool_handle())
# Set graph pool id globally to be able to use symmetric memory
set_graph_pool_id(get_global_graph_memory_pool())
out = self._capture_graph(
graph, get_global_graph_memory_pool(), stream, run_once
)
return graph, out return graph, out
@@ -9,6 +9,10 @@ import triton.language as tl
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
from sglang.srt.model_executor.forward_context import (
get_req_to_token_pool,
get_token_to_kv_pool,
)
class ForwardBatchDeepSeekMHAMixin: class ForwardBatchDeepSeekMHAMixin:
@@ -55,6 +59,7 @@ class ForwardBatchDeepSeekMHAMixin:
def prepare_chunked_kv_indices(self, device: torch.device): def prepare_chunked_kv_indices(self, device: torch.device):
self.prefix_chunk_kv_indices = [] self.prefix_chunk_kv_indices = []
req_to_token = get_req_to_token_pool().req_to_token
for idx in range(self.num_prefix_chunks): for idx in range(self.num_prefix_chunks):
chunk_starts = self.prefix_chunk_starts[idx] chunk_starts = self.prefix_chunk_starts[idx]
chunk_seq_lens = self.prefix_chunk_seq_lens[idx] chunk_seq_lens = self.prefix_chunk_seq_lens[idx]
@@ -66,13 +71,13 @@ class ForwardBatchDeepSeekMHAMixin:
) )
create_chunked_prefix_cache_kv_indices[(self.batch_size,)]( create_chunked_prefix_cache_kv_indices[(self.batch_size,)](
self.req_to_token_pool.req_to_token, req_to_token,
self.req_pool_indices, self.req_pool_indices,
chunk_starts, chunk_starts,
chunk_seq_lens, chunk_seq_lens,
chunk_cu_seq_lens, chunk_cu_seq_lens,
chunk_kv_indices, chunk_kv_indices,
self.req_to_token_pool.req_to_token.shape[1], req_to_token.shape[1],
) )
self.prefix_chunk_kv_indices.append(chunk_kv_indices) self.prefix_chunk_kv_indices.append(chunk_kv_indices)
@@ -111,7 +116,7 @@ class ForwardBatchDeepSeekMHAMixin:
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool
assert isinstance( assert isinstance(
self.token_to_kv_pool, MLATokenToKVPool get_token_to_kv_pool(), MLATokenToKVPool
), "Currently chunked prefix cache can only be used by Deepseek models" ), "Currently chunked prefix cache can only be used by Deepseek models"
if not any(self.extend_prefix_lens_cpu): if not any(self.extend_prefix_lens_cpu):
@@ -191,14 +196,15 @@ class ForwardBatchDeepSeekMHAMixin:
device=self.req_pool_indices.device, device=self.req_pool_indices.device,
) )
kv_indptr[1:] = torch.cumsum(self.seq_lens, dim=0) kv_indptr[1:] = torch.cumsum(self.seq_lens, dim=0)
req_to_token = get_req_to_token_pool().req_to_token
create_flashinfer_kv_indices_triton[(self.batch_size,)]( create_flashinfer_kv_indices_triton[(self.batch_size,)](
self.req_to_token_pool.req_to_token, req_to_token,
self.req_pool_indices, self.req_pool_indices,
self.seq_lens, self.seq_lens,
kv_indptr, kv_indptr,
None, None,
kv_indices, kv_indices,
self.req_to_token_pool.req_to_token.shape[1], req_to_token.shape[1],
) )
self.mha_one_shot_kv_indices = kv_indices self.mha_one_shot_kv_indices = kv_indices
return kv_indices return kv_indices
@@ -63,11 +63,8 @@ from sglang.srt.utils import (
from sglang.srt.utils.common import ceil_align from sglang.srt.utils.common import ceil_align
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator
from sglang.srt.managers.schedule_batch import MultimodalInputs, ScheduleBatch from sglang.srt.managers.schedule_batch import MultimodalInputs, ScheduleBatch
from sglang.srt.mem_cache.memory_pool import KVCache, ReqToTokenPool
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.speculative.spec_info import SpecInput, SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpecInput, SpeculativeAlgorithm
@@ -369,11 +366,6 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# Sampling info # Sampling info
sampling_info: SamplingBatchInfo = None sampling_info: SamplingBatchInfo = None
# Attention backend
req_to_token_pool: ReqToTokenPool = None
token_to_kv_pool: KVCache = None
attn_backend: AttentionBackend = None
# For DP attention # For DP attention
original_global_num_tokens_cpu: Optional[List[int]] = None original_global_num_tokens_cpu: Optional[List[int]] = None
global_num_tokens_cpu: Optional[List[int]] = None global_num_tokens_cpu: Optional[List[int]] = None
@@ -432,9 +424,6 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# Whether to return pooled hidden states (pre-head transformer output) # Whether to return pooled hidden states (pre-head transformer output)
return_pooled_hidden_states: bool = False return_pooled_hidden_states: bool = False
# For hisparse
hisparse_coordinator: Optional[HiSparseCoordinator] = None
# For ngram embedding # For ngram embedding
ngram_embedding_info: Optional[NgramEmbeddingInfo] = None ngram_embedding_info: Optional[NgramEmbeddingInfo] = None
@@ -536,9 +525,6 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
multi_item_delimiter_indices=batch.multi_item_delimiter_indices, multi_item_delimiter_indices=batch.multi_item_delimiter_indices,
lora_ids=[req.lora_id for req in batch.reqs], lora_ids=[req.lora_id for req in batch.reqs],
sampling_info=batch.sampling_info, sampling_info=batch.sampling_info,
req_to_token_pool=model_runner.req_to_token_pool,
token_to_kv_pool=model_runner.token_to_kv_pool,
attn_backend=model_runner.attn_backend,
spec_algorithm=batch.spec_algorithm, spec_algorithm=batch.spec_algorithm,
spec_info=batch.spec_info, spec_info=batch.spec_info,
capture_hidden_mode=capture_hidden_mode, capture_hidden_mode=capture_hidden_mode,
@@ -0,0 +1,84 @@
"""Per-forward-call control context.
Owns ``ForwardContext`` — a frozen dataclass holding control configs the model
layer reads at depth via ``get_forward_context()``. The only mandatory field
today is ``attn_backend``; pool refs are derived from ``attn_backend.*``
(every backend caches them at ``__init__``), so a published ``ForwardContext``
is enough to resolve the active pools without a separate global.
``ModelRunner._forward_raw`` publishes a fresh ``ForwardContext`` for the
duration of each forward; callers that need a per-call override (PDmux
per-stream backend, frozen-KV MTP draft loop, TBO per-child dispatch) use
``dataclasses.replace`` and wrap the override scope with ``forward_context()``.
Distinct from ``sglang.srt.compilation.piecewise_context_manager.ForwardContext``,
which collects compilation-time refs for the piecewise CUDA graph backend.
Concurrency: ``_current`` is a plain module-level global, not thread-local.
This matches the ``global_server_args`` precedent and is safe because each
forward runs synchronously on a single Python thread per worker process. If
worker threads ever share a process, migrate to ``contextvars.ContextVar``.
"""
from __future__ import annotations
from contextlib import contextmanager
from dataclasses import dataclass
from typing import TYPE_CHECKING, Optional
if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.mem_cache.memory_pool import KVCache, ReqToTokenPool
@dataclass(frozen=True, slots=True)
class ForwardContext:
"""Per-forward-call control configs. Read via ``get_forward_context()``;
extend by adding fields here. Frozen so accidental mutation raises at
write time — use ``dataclasses.replace`` for per-call overrides."""
attn_backend: AttentionBackend
_current: Optional[ForwardContext] = None
def set_forward_context(ctx: Optional[ForwardContext]) -> Optional[ForwardContext]:
"""Set the active context; return the previous one for explicit
save/restore. Prefer the ``forward_context()`` context manager."""
global _current
prev, _current = _current, ctx
return prev
def has_forward_context() -> bool:
return _current is not None
def get_forward_context() -> ForwardContext:
assert _current is not None, (
"no forward context active — call forward_context(...) or set_forward_context(...) "
"before reading get_forward_context()."
)
return _current
def get_attn_backend() -> AttentionBackend:
return get_forward_context().attn_backend
def get_token_to_kv_pool() -> KVCache:
return get_attn_backend().token_to_kv_pool
def get_req_to_token_pool() -> ReqToTokenPool:
return get_attn_backend().req_to_token_pool
@contextmanager
def forward_context(ctx: ForwardContext):
prev = set_forward_context(ctx)
try:
yield
finally:
set_forward_context(prev)
+114 -91
View File
@@ -146,6 +146,11 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardMode, ForwardMode,
PPProxyTensors, PPProxyTensors,
) )
from sglang.srt.model_executor.forward_context import (
ForwardContext,
forward_context,
has_forward_context,
)
from sglang.srt.model_executor.hook_manager import register_forward_hooks from sglang.srt.model_executor.hook_manager import register_forward_hooks
from sglang.srt.model_executor.model_runner_kv_cache_mixin import ( from sglang.srt.model_executor.model_runner_kv_cache_mixin import (
ModelRunnerKVCacheMixin, ModelRunnerKVCacheMixin,
@@ -2638,9 +2643,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
seq_lens_cpu=buffers.seq_lens_cpu, seq_lens_cpu=buffers.seq_lens_cpu,
next_token_logits_buffer=buffers.next_token_logits_buffer, next_token_logits_buffer=buffers.next_token_logits_buffer,
orig_seq_lens=buffers.seq_lens, orig_seq_lens=buffers.seq_lens,
req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool=self.token_to_kv_pool,
attn_backend=self.attn_backend,
out_cache_loc=buffers.out_cache_loc, out_cache_loc=buffers.out_cache_loc,
seq_lens_sum=buffers.seq_lens.sum().item(), seq_lens_sum=buffers.seq_lens.sum().item(),
encoder_lens=buffers.encoder_lens, encoder_lens=buffers.encoder_lens,
@@ -2701,8 +2703,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
torch.get_device_module(self.device).synchronize() torch.get_device_module(self.device).synchronize()
self.tp_group.barrier() self.tp_group.barrier()
with torch.inference_mode(), run_ctx or empty_context(): with forward_context(ForwardContext(attn_backend=self.attn_backend)):
run_once() with torch.inference_mode(), run_ctx or empty_context():
run_once()
def maybe_init_ngram_embedding(self): def maybe_init_ngram_embedding(self):
self.use_ngram_embedding = self.model_config.use_ngram_embedding self.use_ngram_embedding = self.model_config.use_ngram_embedding
@@ -2979,6 +2982,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
pp_proxy_tensors=None, pp_proxy_tensors=None,
) -> Union[LogitsProcessorOutput, PPProxyTensors]: ) -> Union[LogitsProcessorOutput, PPProxyTensors]:
# Set extra arguments # Set extra arguments
pdmux_override = False
if not skip_attn_backend_init: if not skip_attn_backend_init:
if hasattr(self.model, "prepare_forward_batch"): if hasattr(self.model, "prepare_forward_batch"):
# Prepare model-specific attention metadata before planning, # Prepare model-specific attention metadata before planning,
@@ -2986,7 +2990,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.model.prepare_forward_batch(forward_batch) self.model.prepare_forward_batch(forward_batch)
if self.server_args.enable_pdmux: if self.server_args.enable_pdmux:
self.decode_attn_backend.init_forward_metadata(forward_batch) self.decode_attn_backend.init_forward_metadata(forward_batch)
forward_batch.attn_backend = self.decode_attn_backend # PDmux selects a per-stream backend; publish it to model-layer
# readers via the active ForwardContext so RadixAttention etc.
# dispatch against the right backend for this forward.
pdmux_override = True
else: else:
self.attn_backend.init_forward_metadata(forward_batch) self.attn_backend.init_forward_metadata(forward_batch)
# FIXME: add pp_proxy_tensors arg to all models # FIXME: add pp_proxy_tensors arg to all models
@@ -3000,7 +3007,8 @@ class ModelRunner(ModelRunnerKVCacheMixin):
if self.device_timer if self.device_timer
else contextlib.nullcontext() else contextlib.nullcontext()
) )
with ctx:
def _do_forward():
return self.model.forward( return self.model.forward(
forward_batch.input_ids, forward_batch.input_ids,
forward_batch.positions, forward_batch.positions,
@@ -3008,6 +3016,14 @@ class ModelRunner(ModelRunnerKVCacheMixin):
**kwargs, **kwargs,
) )
with ctx:
if pdmux_override:
with forward_context(
ForwardContext(attn_backend=self.decode_attn_backend)
):
return _do_forward()
return _do_forward()
def forward_extend( def forward_extend(
self, self,
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
@@ -3216,93 +3232,100 @@ class ModelRunner(ModelRunnerKVCacheMixin):
reinit_attn_backend: bool = False, reinit_attn_backend: bool = False,
split_forward_count: int = 1, split_forward_count: int = 1,
) -> ModelRunnerOutput: ) -> ModelRunnerOutput:
# Check whether can run cuda graph # Honor an outer-published context (spec workers wrap each per-step
mode_check = ( # draft forward with the i-th child backend); otherwise publish this
forward_batch.forward_mode.is_cpu_graph # runner's own attn_backend for the forward.
if self.device == "cpu" if has_forward_context():
else forward_batch.forward_mode.is_cuda_graph ctx_mgr = contextlib.nullcontext()
) else:
can_run_graph = bool( ctx_mgr = forward_context(ForwardContext(attn_backend=self.attn_backend))
mode_check() with ctx_mgr:
and self.graph_runner mode_check = (
and self.graph_runner.can_run(forward_batch) forward_batch.forward_mode.is_cpu_graph
) if self.device == "cpu"
else forward_batch.forward_mode.is_cuda_graph
# Hisparse coordinator
if (
forward_batch.forward_mode.is_decode()
and self.hisparse_coordinator is not None
):
forward_batch.hisparse_coordinator = self.hisparse_coordinator
self.hisparse_coordinator.wait_for_pending_backup()
self.hisparse_coordinator.num_real_reqs.fill_(forward_batch.batch_size)
# Replay cuda graph if applicable
if can_run_graph:
ret = self.graph_runner.replay(
forward_batch,
skip_attn_backend_init=skip_attn_backend_init,
pp_proxy_tensors=pp_proxy_tensors,
) )
can_run_graph = bool(
mode_check()
and self.graph_runner
and self.graph_runner.can_run(forward_batch)
)
# Hisparse coordinator — backends now read it from self.model_runner.
if (
forward_batch.forward_mode.is_decode()
and self.hisparse_coordinator is not None
):
self.hisparse_coordinator.wait_for_pending_backup()
self.hisparse_coordinator.num_real_reqs.fill_(forward_batch.batch_size)
# Replay cuda graph if applicable
if can_run_graph:
ret = self.graph_runner.replay(
forward_batch,
skip_attn_backend_init=skip_attn_backend_init,
pp_proxy_tensors=pp_proxy_tensors,
)
return ModelRunnerOutput(logits_output=ret, can_run_graph=can_run_graph)
# For MLP sync
if forward_batch.global_num_tokens_cpu is not None:
forward_batch.prepare_mlp_sync_batch(self)
else:
forward_batch.prepare_attn_tp_scatter_input(self)
# Normalize num_token_non_padded to be local to this attention TP rank if needed.
if (
forward_batch.num_token_non_padded is not None
and forward_batch.global_num_tokens_gpu is not None
and require_gathered_buffer(self.server_args)
and not is_dsa_enable_prefill_cp()
):
forward_batch.adjust_num_token_non_padded_for_attn_tp(
server_args=self.server_args,
)
if self.is_hybrid_swa:
self.token_to_kv_pool.invalidate_loc_cache()
# Hisparse coordinator — backends now read it from self.model_runner.
if self.hisparse_coordinator is not None:
self.hisparse_coordinator.num_real_reqs.fill_(forward_batch.batch_size)
# Forward without cuda graph
if forward_batch.forward_mode.is_decode():
ret = self.forward_decode(
forward_batch,
skip_attn_backend_init=skip_attn_backend_init,
pp_proxy_tensors=pp_proxy_tensors,
)
elif forward_batch.forward_mode.is_split_prefill():
ret = self.forward_split_prefill(
forward_batch,
reinit_attn_backend=reinit_attn_backend,
forward_count=split_forward_count,
)
elif forward_batch.forward_mode.is_extend(include_draft_extend_v2=True):
ret, can_run_graph = self.forward_extend(
forward_batch,
skip_attn_backend_init=skip_attn_backend_init,
pp_proxy_tensors=pp_proxy_tensors,
)
elif forward_batch.forward_mode.is_idle():
ret = self.forward_idle(
forward_batch, pp_proxy_tensors=pp_proxy_tensors
)
else:
raise ValueError(f"Invalid forward mode: {forward_batch.forward_mode}")
if (
forward_batch.global_num_tokens_cpu is not None
and self.pp_group.is_last_rank
):
forward_batch.post_forward_mlp_sync_batch(ret)
return ModelRunnerOutput(logits_output=ret, can_run_graph=can_run_graph) return ModelRunnerOutput(logits_output=ret, can_run_graph=can_run_graph)
# For MLP sync
if forward_batch.global_num_tokens_cpu is not None:
forward_batch.prepare_mlp_sync_batch(self)
else:
forward_batch.prepare_attn_tp_scatter_input(self)
# Normalize num_token_non_padded to be local to this attention TP rank if needed.
if (
forward_batch.num_token_non_padded is not None
and forward_batch.global_num_tokens_gpu is not None
and require_gathered_buffer(self.server_args)
and not is_dsa_enable_prefill_cp()
):
forward_batch.adjust_num_token_non_padded_for_attn_tp(
server_args=self.server_args,
)
if self.is_hybrid_swa:
self.token_to_kv_pool.invalidate_loc_cache()
# Hisparse coordinator
forward_batch.hisparse_coordinator = self.hisparse_coordinator
if self.hisparse_coordinator is not None:
self.hisparse_coordinator.num_real_reqs.fill_(forward_batch.batch_size)
# Forward without cuda graph
if forward_batch.forward_mode.is_decode():
ret = self.forward_decode(
forward_batch,
skip_attn_backend_init=skip_attn_backend_init,
pp_proxy_tensors=pp_proxy_tensors,
)
elif forward_batch.forward_mode.is_split_prefill():
ret = self.forward_split_prefill(
forward_batch,
reinit_attn_backend=reinit_attn_backend,
forward_count=split_forward_count,
)
elif forward_batch.forward_mode.is_extend(include_draft_extend_v2=True):
ret, can_run_graph = self.forward_extend(
forward_batch,
skip_attn_backend_init=skip_attn_backend_init,
pp_proxy_tensors=pp_proxy_tensors,
)
elif forward_batch.forward_mode.is_idle():
ret = self.forward_idle(forward_batch, pp_proxy_tensors=pp_proxy_tensors)
else:
raise ValueError(f"Invalid forward mode: {forward_batch.forward_mode}")
if (
forward_batch.global_num_tokens_cpu is not None
and self.pp_group.is_last_rank
):
forward_batch.post_forward_mlp_sync_batch(ret)
return ModelRunnerOutput(logits_output=ret, can_run_graph=can_run_graph)
def _preprocess_logits( def _preprocess_logits(
self, logits_output: LogitsProcessorOutput, sampling_info: SamplingBatchInfo self, logits_output: LogitsProcessorOutput, sampling_info: SamplingBatchInfo
): ):
@@ -58,6 +58,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardMode, ForwardMode,
PPProxyTensors, PPProxyTensors,
) )
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
from sglang.srt.utils import ( from sglang.srt.utils import (
get_available_gpu_memory, get_available_gpu_memory,
@@ -387,9 +388,6 @@ class PiecewiseCudaGraphRunner:
next_token_logits_buffer=None, next_token_logits_buffer=None,
orig_seq_lens=torch.tensor([num_tokens], device=self.device), orig_seq_lens=torch.tensor([num_tokens], device=self.device),
seq_lens_cpu=torch.tensor([num_tokens], device="cpu"), seq_lens_cpu=torch.tensor([num_tokens], device="cpu"),
req_to_token_pool=self.model_runner.req_to_token_pool,
token_to_kv_pool=self.model_runner.token_to_kv_pool,
attn_backend=self.model_runner.attn_backend,
out_cache_loc=out_cache_loc, out_cache_loc=out_cache_loc,
seq_lens_sum=num_tokens, seq_lens_sum=num_tokens,
mamba_track_indices=mamba_track_indices, mamba_track_indices=mamba_track_indices,
@@ -425,18 +423,21 @@ class PiecewiseCudaGraphRunner:
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len(None, num_tokens, forward_batch.dp_padding_mode.is_max_len()) set_dp_buffer_len(None, num_tokens, forward_batch.dp_padding_mode.is_max_len())
set_is_extend_in_batch(False) set_is_extend_in_batch(False)
with set_forward_context( with forward_context(
forward_batch, ForwardContext(attn_backend=self.model_runner.attn_backend)
self.attention_layers,
self.quant_config,
self.moe_layers,
self.moe_fusions,
): ):
_ = self.model_runner.model.forward( with set_forward_context(
forward_batch.input_ids,
forward_batch.positions,
forward_batch, forward_batch,
) self.attention_layers,
self.quant_config,
self.moe_layers,
self.moe_fusions,
):
_ = self.model_runner.model.forward(
forward_batch.input_ids,
forward_batch.positions,
forward_batch,
)
def _cache_loc_dtype(self): def _cache_loc_dtype(self):
return torch.int64 if not is_npu() else torch.int32 return torch.int64 if not is_npu() else torch.int32
@@ -554,9 +555,6 @@ class PiecewiseCudaGraphRunner:
next_token_logits_buffer=None, next_token_logits_buffer=None,
orig_seq_lens=torch.tensor([num_tokens], device=self.device), orig_seq_lens=torch.tensor([num_tokens], device=self.device),
seq_lens_cpu=torch.tensor([num_tokens], device="cpu"), seq_lens_cpu=torch.tensor([num_tokens], device="cpu"),
req_to_token_pool=self.model_runner.req_to_token_pool,
token_to_kv_pool=self.model_runner.token_to_kv_pool,
attn_backend=self.model_runner.attn_backend,
out_cache_loc=out_cache_loc, out_cache_loc=out_cache_loc,
seq_lens_sum=num_tokens, seq_lens_sum=num_tokens,
mamba_track_indices=mamba_track_indices, mamba_track_indices=mamba_track_indices,
@@ -586,52 +584,59 @@ class PiecewiseCudaGraphRunner:
lora_ids=None, lora_ids=None,
return_pooled_hidden_states=self.capture_return_pooled_hidden_states, return_pooled_hidden_states=self.capture_return_pooled_hidden_states,
) )
# Setup hooks below read get_attn_backend() and must run inside the
# same ForwardContext as the warmup/capture forward.
with forward_context(
ForwardContext(attn_backend=self.model_runner.attn_backend)
):
self.tbo_plugin.capture_one_batch_size(forward_batch, num_tokens=num_tokens) self.tbo_plugin.capture_one_batch_size(forward_batch, num_tokens=num_tokens)
if lora_ids is not None: if lora_ids is not None:
self.model_runner.lora_manager.prepare_lora_batch(forward_batch) self.model_runner.lora_manager.prepare_lora_batch(forward_batch)
self.model_runner.attn_backend.init_forward_metadata(forward_batch) self.model_runner.attn_backend.init_forward_metadata(forward_batch)
# Run and capture # Run and capture
def run_once(): def run_once():
# Invalidate SWA loc cache — same fix as in cuda_graph_runner.run_once. # Invalidate SWA loc cache — same fix as in cuda_graph_runner.run_once.
if self.model_runner.is_hybrid_swa: if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache() self.model_runner.token_to_kv_pool.invalidate_loc_cache()
# Clean intermediate result cache for DP attention # Clean intermediate result cache for DP attention
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = (
set_dp_buffer_len( None
global_dp_buffer_len,
num_tokens,
forward_batch.dp_padding_mode.is_max_len(),
)
# FIXME: the implementation is hacky. `is_extend_in_batch`` is for determining the deepep mode.
# It is True in this context but we need to set it to use low latency deepep mode.
set_is_extend_in_batch(False)
kwargs = {}
with set_forward_context(
forward_batch,
self.attention_layers,
self.quant_config,
self.moe_layers,
self.moe_fusions,
):
self.model_runner.model.forward(
forward_batch.input_ids,
forward_batch.positions,
forward_batch,
**kwargs,
) )
return set_dp_buffer_len(
global_dp_buffer_len,
num_tokens,
forward_batch.dp_padding_mode.is_max_len(),
)
# FIXME: the implementation is hacky. `is_extend_in_batch`` is for determining the deepep mode.
# It is True in this context but we need to set it to use low latency deepep mode.
set_is_extend_in_batch(False)
# run twice for warmup at the first time and cuda graph capture at the second time kwargs = {}
# detail lies in sglang/python/sglang/srt/compilation/cuda_piecewise_backend.py with set_forward_context(
for _ in range(2): forward_batch,
self.device_module.synchronize() self.attention_layers,
self.model_runner.tp_group.barrier() self.quant_config,
run_once() self.moe_layers,
self.moe_fusions,
):
self.model_runner.model.forward(
forward_batch.input_ids,
forward_batch.positions,
forward_batch,
**kwargs,
)
return
# run twice for warmup at the first time and cuda graph capture at the second time
# detail lies in sglang/python/sglang/srt/compilation/cuda_piecewise_backend.py
for _ in range(2):
self.device_module.synchronize()
self.model_runner.tp_group.barrier()
run_once()
return return
@@ -733,9 +738,6 @@ class PiecewiseCudaGraphRunner:
next_token_logits_buffer=next_token_logits_buffer, next_token_logits_buffer=next_token_logits_buffer,
orig_seq_lens=forward_batch.orig_seq_lens, orig_seq_lens=forward_batch.orig_seq_lens,
seq_lens_cpu=forward_batch.seq_lens_cpu, seq_lens_cpu=forward_batch.seq_lens_cpu,
req_to_token_pool=self.model_runner.req_to_token_pool,
token_to_kv_pool=self.model_runner.token_to_kv_pool,
attn_backend=self.model_runner.attn_backend,
out_cache_loc=out_cache_loc, out_cache_loc=out_cache_loc,
seq_lens_sum=forward_batch.seq_lens_sum, seq_lens_sum=forward_batch.seq_lens_sum,
mamba_track_indices=mamba_track_indices, mamba_track_indices=mamba_track_indices,
@@ -1,5 +1,6 @@
from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods import ( from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods import (
AttnForwardMethod, AttnForwardMethod,
) )
@@ -153,7 +154,7 @@ def handle_attention_dsa(attn, forward_batch):
in init_forward_metadata. Read the decision from backend.use_mha. in init_forward_metadata. Read the decision from backend.use_mha.
""" """
backend = forward_batch.attn_backend backend = get_attn_backend()
if isinstance(backend, TboAttnBackend): # if enable tbo, get primary backend if isinstance(backend, TboAttnBackend): # if enable tbo, get primary backend
backend = backend.primary backend = backend.primary
if hasattr(backend, "use_mha") and backend.use_mha: if hasattr(backend, "use_mha") and backend.use_mha:
@@ -10,6 +10,10 @@ from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
from sglang.srt.layers.attention.utils import concat_and_cast_mha_k_triton from sglang.srt.layers.attention.utils import concat_and_cast_mha_k_triton
from sglang.srt.layers.communicator import get_attn_tp_context from sglang.srt.layers.communicator import get_attn_tp_context
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import (
get_attn_backend,
get_token_to_kv_pool,
)
from sglang.srt.models.deepseek_common.utils import ( from sglang.srt.models.deepseek_common.utils import (
_is_cuda, _is_cuda,
_is_hip, _is_hip,
@@ -38,7 +42,7 @@ if _use_aiter_gfx95:
def _resolve_attn_backend(forward_batch: ForwardBatch): def _resolve_attn_backend(forward_batch: ForwardBatch):
backend = forward_batch.attn_backend backend = get_attn_backend()
if isinstance(backend, TboAttnBackend): if isinstance(backend, TboAttnBackend):
backend = backend.primary backend = backend.primary
return backend return backend
@@ -334,8 +338,8 @@ class DeepseekMHAForwardMixin:
# Only initialize the info once # Only initialize the info once
if has_extend_prefix and forward_batch.num_prefix_chunks is None: if has_extend_prefix and forward_batch.num_prefix_chunks is None:
forward_batch.prepare_chunked_prefix_cache_info(q.device) forward_batch.prepare_chunked_prefix_cache_info(q.device)
if hasattr(forward_batch.attn_backend, "init_mha_chunk_metadata"): if hasattr(get_attn_backend(), "init_mha_chunk_metadata"):
forward_batch.attn_backend.init_mha_chunk_metadata(forward_batch) get_attn_backend().init_mha_chunk_metadata(forward_batch)
forward_batch.mha_return_lse = has_extend_prefix forward_batch.mha_return_lse = has_extend_prefix
# Do mha for extended part without prefix # Do mha for extended part without prefix
@@ -380,8 +384,8 @@ class DeepseekMHAForwardMixin:
# Only initialize the info once # Only initialize the info once
if has_extend_prefix and forward_batch.num_prefix_chunks is None: if has_extend_prefix and forward_batch.num_prefix_chunks is None:
forward_batch.num_prefix_chunks = 0 forward_batch.num_prefix_chunks = 0
if hasattr(forward_batch.attn_backend, "init_mha_chunk_metadata"): if hasattr(get_attn_backend(), "init_mha_chunk_metadata"):
forward_batch.attn_backend.init_mha_chunk_metadata(forward_batch) get_attn_backend().init_mha_chunk_metadata(forward_batch)
forward_batch.mha_return_lse = False forward_batch.mha_return_lse = False
# Do mha for extended part without prefix # Do mha for extended part without prefix
forward_batch.set_attn_attend_prefix_cache(False) forward_batch.set_attn_attend_prefix_cache(False)
@@ -449,12 +453,12 @@ class DeepseekMHAForwardMixin:
): ):
if _is_cuda or _use_aiter_gfx95: if _is_cuda or _use_aiter_gfx95:
# Save latent cache # Save latent cache
forward_batch.token_to_kv_pool.set_mla_kv_buffer( get_token_to_kv_pool().set_mla_kv_buffer(
self.attn_mha, forward_batch.out_cache_loc, kv_a.unsqueeze(1), k_pe self.attn_mha, forward_batch.out_cache_loc, kv_a.unsqueeze(1), k_pe
) )
elif _is_npu: elif _is_npu:
# To reduce a time-costing split operation # To reduce a time-costing split operation
forward_batch.token_to_kv_pool.set_kv_buffer( get_token_to_kv_pool().set_kv_buffer(
self.attn_mha, forward_batch.out_cache_loc, kv_a.unsqueeze(1), k_pe self.attn_mha, forward_batch.out_cache_loc, kv_a.unsqueeze(1), k_pe
) )
else: else:
@@ -462,7 +466,7 @@ class DeepseekMHAForwardMixin:
latent_cache[:, :, self.kv_lora_rank :] = k_pe.clone() latent_cache[:, :, self.kv_lora_rank :] = k_pe.clone()
# Save latent cache # Save latent cache
forward_batch.token_to_kv_pool.set_kv_buffer( get_token_to_kv_pool().set_kv_buffer(
self.attn_mha, forward_batch.out_cache_loc, latent_cache, None self.attn_mha, forward_batch.out_cache_loc, latent_cache, None
) )
@@ -473,12 +477,12 @@ class DeepseekMHAForwardMixin:
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
): ):
if _is_cuda or _use_aiter_gfx95: if _is_cuda or _use_aiter_gfx95:
kv_a, k_pe = forward_batch.token_to_kv_pool.get_mla_kv_buffer( kv_a, k_pe = get_token_to_kv_pool().get_mla_kv_buffer(
self.attn_mha, kv_indices, dst_dtype self.attn_mha, kv_indices, dst_dtype
) )
kv_a = kv_a.squeeze(1) kv_a = kv_a.squeeze(1)
else: else:
latent_cache_buf = forward_batch.token_to_kv_pool.get_key_buffer( latent_cache_buf = get_token_to_kv_pool().get_key_buffer(
self.attn_mha.layer_id self.attn_mha.layer_id
) )
latent_cache = latent_cache_buf[kv_indices].contiguous().to(dst_dtype) latent_cache = latent_cache_buf[kv_indices].contiguous().to(dst_dtype)
@@ -498,7 +502,7 @@ class DeepseekMHAForwardMixin:
Returns: (kv_a, k_pe) both in BF16 Returns: (kv_a, k_pe) both in BF16
""" """
backend = forward_batch.attn_backend backend = get_attn_backend()
if isinstance(backend, TboAttnBackend): # if enable tbo, get primary backend if isinstance(backend, TboAttnBackend): # if enable tbo, get primary backend
backend = backend.primary backend = backend.primary
kv_indices = backend.forward_metadata.page_table_1_flattened kv_indices = backend.forward_metadata.page_table_1_flattened
@@ -506,9 +510,7 @@ class DeepseekMHAForwardMixin:
kv_indices is not None kv_indices is not None
), "page_table_1_flattened should have been generated for FP8 MHA path" ), "page_table_1_flattened should have been generated for FP8 MHA path"
kv_cache_fp8 = forward_batch.token_to_kv_pool.get_key_buffer( kv_cache_fp8 = get_token_to_kv_pool().get_key_buffer(self.attn_mha.layer_id)
self.attn_mha.layer_id
)
kv_latent_bf16 = dequantize_k_cache_paged(kv_cache_fp8, kv_indices) kv_latent_bf16 = dequantize_k_cache_paged(kv_cache_fp8, kv_indices)
@@ -544,7 +546,7 @@ class DeepseekMHAForwardMixin:
self.current_attention_backend == "fa3" self.current_attention_backend == "fa3"
and self.kv_cache_dtype != "auto" and self.kv_cache_dtype != "auto"
): ):
attn_dtype = forward_batch.token_to_kv_pool.dtype attn_dtype = get_token_to_kv_pool().dtype
else: else:
attn_dtype = k_nope.dtype attn_dtype = k_nope.dtype
k = k_nope.new_empty(*k_shape, dtype=attn_dtype) k = k_nope.new_empty(*k_shape, dtype=attn_dtype)
@@ -23,6 +23,10 @@ from sglang.srt.lora.deepseek_mla_correction import (
is_kv_b_lora_active, is_kv_b_lora_active,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import (
get_attn_backend,
get_token_to_kv_pool,
)
from sglang.srt.models.deepseek_common.utils import ( from sglang.srt.models.deepseek_common.utils import (
FORWARD_ABSORB_CORE_ATTENTION_BACKENDS, FORWARD_ABSORB_CORE_ATTENTION_BACKENDS,
_is_cpu, _is_cpu,
@@ -420,9 +424,7 @@ class DeepseekMLAForwardMixin:
q_pe, q_pe,
k_nope, k_nope,
k_pe, k_pe,
forward_batch.token_to_kv_pool.get_key_buffer( get_token_to_kv_pool().get_key_buffer(self.attn_mqa.layer_id),
self.attn_mqa.layer_id
),
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
positions, positions,
cos, cos,
@@ -516,9 +518,7 @@ class DeepseekMLAForwardMixin:
q_pe, q_pe,
k_nope, k_nope,
k_pe, k_pe,
forward_batch.token_to_kv_pool.get_key_buffer( get_token_to_kv_pool().get_key_buffer(self.attn_mqa.layer_id),
self.attn_mqa.layer_id
),
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
positions, positions,
cos, cos,
@@ -694,7 +694,7 @@ class DeepseekMLAForwardMixin:
return ( return (
get_global_server_args().dsa_decode_backend == "trtllm" get_global_server_args().dsa_decode_backend == "trtllm"
or get_global_server_args().dsa_prefill_backend == "trtllm" or get_global_server_args().dsa_prefill_backend == "trtllm"
) and forward_batch.attn_backend.kv_cache_dtype == torch.float8_e4m3fn ) and get_attn_backend().kv_cache_dtype == torch.float8_e4m3fn
return ( return (
self.current_attention_backend in ("trtllm_mla", "tokenspeed_mla") self.current_attention_backend in ("trtllm_mla", "tokenspeed_mla")
@@ -702,7 +702,7 @@ class DeepseekMLAForwardMixin:
forward_batch.forward_mode.is_decode_or_idle() forward_batch.forward_mode.is_decode_or_idle()
or forward_batch.forward_mode.is_target_verify() or forward_batch.forward_mode.is_target_verify()
) )
and forward_batch.attn_backend.data_type == torch.float8_e4m3fn and get_attn_backend().data_type == torch.float8_e4m3fn
) )
def _skip_rope_for_dsa_tilelang_fused(self: DeepseekV2AttentionMLA) -> bool: def _skip_rope_for_dsa_tilelang_fused(self: DeepseekV2AttentionMLA) -> bool:
@@ -7,6 +7,10 @@ import torch
from sglang.srt.layers.quantization.fp8_kernel import per_tensor_quant_mla_fp8 from sglang.srt.layers.quantization.fp8_kernel import per_tensor_quant_mla_fp8
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import (
get_attn_backend,
get_token_to_kv_pool,
)
from sglang.srt.models.deepseek_common.utils import ( from sglang.srt.models.deepseek_common.utils import (
_is_cuda, _is_cuda,
_is_hip, _is_hip,
@@ -108,10 +112,10 @@ class DeepseekMLARocmForwardMixin:
device=q.device, device=q.device,
) )
attn_logits, _, kv_indptr, kv_indices, _, _, _ = ( attn_logits, _, kv_indptr, kv_indices, _, _, _ = (
forward_batch.attn_backend.forward_metadata get_attn_backend().forward_metadata
) )
cos_sin_cache = self.rotary_emb.cos_sin_cache cos_sin_cache = self.rotary_emb.cos_sin_cache
num_kv_split = forward_batch.attn_backend.num_kv_splits num_kv_split = get_attn_backend().num_kv_splits
sm_scale = self.attn_mqa.scaling sm_scale = self.attn_mqa.scaling
if attn_logits is None: if attn_logits is None:
attn_logits = torch.empty( attn_logits = torch.empty(
@@ -126,12 +130,10 @@ class DeepseekMLARocmForwardMixin:
) )
# save current latent cache. # save current latent cache.
forward_batch.token_to_kv_pool.set_kv_buffer( get_token_to_kv_pool().set_kv_buffer(
self.attn_mqa, forward_batch.out_cache_loc, k_input, None self.attn_mqa, forward_batch.out_cache_loc, k_input, None
) )
key_cache_buf = forward_batch.token_to_kv_pool.get_key_buffer( key_cache_buf = get_token_to_kv_pool().get_key_buffer(self.attn_mqa.layer_id)
self.attn_mqa.layer_id
)
val_cache_buf = key_cache_buf[..., : self.kv_lora_rank] val_cache_buf = key_cache_buf[..., : self.kv_lora_rank]
return ( return (
@@ -194,7 +196,7 @@ class DeepseekMLARocmForwardMixin:
if enable_rope_fusion: if enable_rope_fusion:
k_input[..., self.kv_lora_rank :] = k_pe_output k_input[..., self.kv_lora_rank :] = k_pe_output
forward_batch.token_to_kv_pool.set_kv_buffer( get_token_to_kv_pool().set_kv_buffer(
self.attn_mqa, forward_batch.out_cache_loc, k_input, None self.attn_mqa, forward_batch.out_cache_loc, k_input, None
) )
+17 -8
View File
@@ -78,6 +78,10 @@ from sglang.srt.model_executor.cuda_graph_runner import (
get_is_capture_mode, get_is_capture_mode,
) )
from sglang.srt.model_executor.forward_batch_info import PPProxyTensors from sglang.srt.model_executor.forward_batch_info import PPProxyTensors
from sglang.srt.model_executor.forward_context import (
get_attn_backend,
get_token_to_kv_pool,
)
from sglang.srt.model_loader.utils import maybe_executor_submit, should_async_load from sglang.srt.model_loader.utils import maybe_executor_submit, should_async_load
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.dbrx import ReplicatedLinear from sglang.srt.models.dbrx import ReplicatedLinear
@@ -398,7 +402,7 @@ class MQALayer(nn.Module):
kv = qkv_a[..., self.q_lora_rank :] kv = qkv_a[..., self.q_lora_rank :]
else: else:
kv, _ = self.wkv(x) kv, _ = self.wkv(x)
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = get_token_to_kv_pool()
if TYPE_CHECKING: if TYPE_CHECKING:
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
token_to_kv_pool.set_swa_key_buffer_radix_fused_norm_rope( token_to_kv_pool.set_swa_key_buffer_radix_fused_norm_rope(
@@ -467,6 +471,7 @@ class MQALayer(nn.Module):
x=x, x=x,
q_lora=q_lora, q_lora=q_lora,
forward_batch=forward_batch, forward_batch=forward_batch,
attn_backend=attn_backend,
enable_multi_stream=True, enable_multi_stream=True,
q_lora_ready=q_lora_ready, q_lora_ready=q_lora_ready,
) )
@@ -533,7 +538,12 @@ class MQALayer(nn.Module):
del qkv_a del qkv_a
if self.indexer is not None: if self.indexer is not None:
self.indexer(x=x, q_lora=q_lora, forward_batch=forward_batch) self.indexer(
x=x,
q_lora=q_lora,
forward_batch=forward_batch,
attn_backend=attn_backend,
)
if self.compressor is not None: if self.compressor is not None:
attn_backend.forward_core_compressor( attn_backend.forward_core_compressor(
x, x,
@@ -556,7 +566,7 @@ class MQALayer(nn.Module):
), "short-circuiting allreduce will lead to hangs" ), "short-circuiting allreduce will lead to hangs"
return x return x
attn_backend = forward_batch.attn_backend attn_backend = get_attn_backend()
if TYPE_CHECKING: if TYPE_CHECKING:
assert isinstance( assert isinstance(
attn_backend, attn_backend,
@@ -1130,7 +1140,7 @@ class DeepseekV4Model(nn.Module):
# Upgrade lazy raw metadata on the main stream once before any layer # Upgrade lazy raw metadata on the main stream once before any layer
# forks alt-streams; later per-layer calls become no-ops. # forks alt-streams; later per-layer calls become no-ops.
forward_batch.attn_backend._maybe_upgrade_forward_metadata() get_attn_backend()._maybe_upgrade_forward_metadata()
for i in range(self.start_layer, self.end_layer): for i in range(self.start_layer, self.end_layer):
layer = self.layers[i] layer = self.layers[i]
@@ -1278,15 +1288,14 @@ class DeepseekV4ForCausalLM(nn.Module):
forward_batch.seq_lens_cpu.tolist(), forward_batch.seq_lens_cpu.tolist(),
) )
if is_dsa_prefill_cp_round_robin_split(): if is_dsa_prefill_cp_round_robin_split():
metadata = forward_batch.attn_backend.forward_metadata attn_backend = get_attn_backend()
metadata = attn_backend.forward_metadata
core_meta = metadata.core_attn_metadata core_meta = metadata.core_attn_metadata
core_meta.apply_cp_reindex() core_meta.apply_cp_reindex()
core_meta.init_flashmla_related() core_meta.init_flashmla_related()
if metadata.indexer_metadata is not None: if metadata.indexer_metadata is not None:
metadata.indexer_metadata = ( metadata.indexer_metadata = (
forward_batch.attn_backend.init_forward_metadata_indexer( attn_backend.init_forward_metadata_indexer(core_meta)
core_meta
)
) )
with get_attn_tp_context().maybe_input_scattered(forward_batch): with get_attn_tp_context().maybe_input_scattered(forward_batch):
@@ -37,6 +37,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
VocabParallelEmbedding, VocabParallelEmbedding,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.models.deepseek_v4 import DeepseekV4DecoderLayer, DeepseekV4ForCausalLM from sglang.srt.models.deepseek_v4 import DeepseekV4DecoderLayer, DeepseekV4ForCausalLM
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
@@ -250,15 +251,14 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM):
forward_batch.seq_lens_cpu.tolist(), forward_batch.seq_lens_cpu.tolist(),
) )
if is_dsa_prefill_cp_round_robin_split(): if is_dsa_prefill_cp_round_robin_split():
metadata = forward_batch.attn_backend.forward_metadata attn_backend = get_attn_backend()
metadata = attn_backend.forward_metadata
core_meta = metadata.core_attn_metadata core_meta = metadata.core_attn_metadata
core_meta.apply_cp_reindex() core_meta.apply_cp_reindex()
core_meta.init_flashmla_related() core_meta.init_flashmla_related()
if metadata.indexer_metadata is not None: if metadata.indexer_metadata is not None:
metadata.indexer_metadata = ( metadata.indexer_metadata = (
forward_batch.attn_backend.init_forward_metadata_indexer( attn_backend.init_forward_metadata_indexer(core_meta)
core_meta
)
) )
hidden_states, pre_hc_head = self.model(input_ids, positions, forward_batch) hidden_states, pre_hc_head = self.model(input_ids, positions, forward_batch)
+2 -1
View File
@@ -33,6 +33,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
VocabParallelEmbedding, VocabParallelEmbedding,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_cuda, make_layers from sglang.srt.utils import add_prefix, is_cuda, make_layers
@@ -338,7 +339,7 @@ class FalconH1HybridAttentionDecoderLayer(nn.Module):
) )
attention_hidden_states = attention_hidden_states * self.attn_out_multiplier attention_hidden_states = attention_hidden_states * self.attn_out_multiplier
attn_backend = forward_batch.attn_backend attn_backend = get_attn_backend()
assert isinstance(attn_backend, HybridLinearAttnBackend) assert isinstance(attn_backend, HybridLinearAttnBackend)
assert isinstance(attn_backend.linear_attn_backend, Mamba2AttnBackend) assert isinstance(attn_backend.linear_attn_backend, Mamba2AttnBackend)
# Mamba block # Mamba block
+4 -3
View File
@@ -40,6 +40,7 @@ from sglang.srt.managers.schedule_batch import (
flatten_nested_list, flatten_nested_list,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.model_loader.weight_utils import ( from sglang.srt.model_loader.weight_utils import (
default_weight_loader, default_weight_loader,
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
@@ -220,7 +221,7 @@ class Gemma3ForConditionalGeneration(PreTrainedModel):
mask_dtype: torch.dtype, mask_dtype: torch.dtype,
): ):
"""Prepare attention masks for multimodal inputs.""" """Prepare attention masks for multimodal inputs."""
if isinstance(forward_batch.attn_backend, TritonAttnBackend): if isinstance(get_attn_backend(), TritonAttnBackend):
assert forward_batch.forward_mode == ForwardMode.EXTEND assert forward_batch.forward_mode == ForwardMode.EXTEND
bidirectional_attn_masks_list = [] bidirectional_attn_masks_list = []
bidirectional_attn_mask_indptr = torch.zeros( bidirectional_attn_mask_indptr = torch.zeros(
@@ -265,10 +266,10 @@ class Gemma3ForConditionalGeneration(PreTrainedModel):
bidirectional_attn_masks = torch.cat( bidirectional_attn_masks = torch.cat(
bidirectional_attn_masks_list, dim=0 bidirectional_attn_masks_list, dim=0
) )
forward_batch.attn_backend.forward_metadata.mask_indptr = ( get_attn_backend().forward_metadata.mask_indptr = (
bidirectional_attn_mask_indptr bidirectional_attn_mask_indptr
) )
forward_batch.attn_backend.forward_metadata.custom_mask = ( get_attn_backend().forward_metadata.custom_mask = (
bidirectional_attn_masks bidirectional_attn_masks
) )
+4 -5
View File
@@ -52,6 +52,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardMode, ForwardMode,
PPProxyTensors, PPProxyTensors,
) )
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.model_loader.weight_utils import ( from sglang.srt.model_loader.weight_utils import (
default_weight_loader, default_weight_loader,
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
@@ -315,7 +316,7 @@ class Gemma4ForConditionalGeneration(PreTrainedModel):
TODO(kpham-sgl): Guard appropriately for gemma3_mm.py:prepare_attn_masks() TODO(kpham-sgl): Guard appropriately for gemma3_mm.py:prepare_attn_masks()
""" """
if not isinstance(forward_batch.attn_backend, TritonAttnBackend): if not isinstance(get_attn_backend(), TritonAttnBackend):
logger.warning_once( logger.warning_once(
"Bidirectional attention for image tokens requires TritonAttnBackend. " "Bidirectional attention for image tokens requires TritonAttnBackend. "
"Falling back to causal attention, which may degrade image quality." "Falling back to causal attention, which may degrade image quality."
@@ -389,12 +390,10 @@ class Gemma4ForConditionalGeneration(PreTrainedModel):
) )
if bidirectional_attn_masks_list: if bidirectional_attn_masks_list:
bidirectional_attn_masks = torch.cat(bidirectional_attn_masks_list, dim=0) bidirectional_attn_masks = torch.cat(bidirectional_attn_masks_list, dim=0)
forward_batch.attn_backend.forward_metadata.mask_indptr = ( get_attn_backend().forward_metadata.mask_indptr = (
bidirectional_attn_mask_indptr bidirectional_attn_mask_indptr
) )
forward_batch.attn_backend.forward_metadata.custom_mask = ( get_attn_backend().forward_metadata.custom_mask = bidirectional_attn_masks
bidirectional_attn_masks
)
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor: def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
vt = self.vision_tower vt = self.vision_tower
+2 -1
View File
@@ -29,6 +29,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
VocabParallelEmbedding, VocabParallelEmbedding,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.transformers import maybe_prefix from sglang.srt.models.transformers import maybe_prefix
from sglang.srt.utils import make_layers from sglang.srt.utils import make_layers
@@ -139,7 +140,7 @@ class GraniteMoeHybridMambaDecoderLayer(nn.Module):
hidden_states = self.input_layernorm(hidden_states) hidden_states = self.input_layernorm(hidden_states)
output = torch.empty_like(hidden_states) output = torch.empty_like(hidden_states)
attn_backend = forward_batch.attn_backend attn_backend = get_attn_backend()
assert isinstance(attn_backend, HybridLinearAttnBackend) assert isinstance(attn_backend, HybridLinearAttnBackend)
assert isinstance(attn_backend.linear_attn_backend, Mamba2AttnBackend) assert isinstance(attn_backend.linear_attn_backend, Mamba2AttnBackend)
attn_backend.linear_attn_backend.forward( attn_backend.linear_attn_backend.forward(
+4 -5
View File
@@ -28,6 +28,7 @@ from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.rotary_embedding import get_rope from sglang.srt.layers.rotary_embedding import get_rope
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.qwen2 import Qwen2MLP, Qwen2Model from sglang.srt.models.qwen2 import Qwen2MLP, Qwen2Model
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
@@ -258,11 +259,9 @@ class JetBlock(nn.Module):
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
) -> torch.Tensor: ) -> torch.Tensor:
assert isinstance(forward_batch.attn_backend, HybridLinearAttnBackend) assert isinstance(get_attn_backend(), HybridLinearAttnBackend)
assert isinstance( assert isinstance(get_attn_backend().linear_attn_backend, MambaAttnBackendBase)
forward_batch.attn_backend.linear_attn_backend, MambaAttnBackendBase linear_attn_backend = get_attn_backend().linear_attn_backend
)
linear_attn_backend = forward_batch.attn_backend.linear_attn_backend
forward_metadata = linear_attn_backend.forward_metadata forward_metadata = linear_attn_backend.forward_metadata
layer_cache = linear_attn_backend.req_to_token_pool.mamba2_layer_cache( layer_cache = linear_attn_backend.req_to_token_pool.mamba2_layer_cache(
self.layer_id self.layer_id
+3 -4
View File
@@ -40,6 +40,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
VocabParallelEmbedding, VocabParallelEmbedding,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_req_to_token_pool
from sglang.srt.model_loader.weight_utils import ( from sglang.srt.model_loader.weight_utils import (
default_weight_loader, default_weight_loader,
sharded_weight_loader, sharded_weight_loader,
@@ -263,12 +264,10 @@ class Lfm2ShortConv(nn.Module):
if forward_batch.forward_mode.is_idle(): if forward_batch.forward_mode.is_idle():
return hidden_states return hidden_states
layer_cache = forward_batch.req_to_token_pool.mamba2_layer_cache(self.layer_idx) layer_cache = get_req_to_token_pool().mamba2_layer_cache(self.layer_idx)
conv_state = layer_cache.conv[0] conv_state = layer_cache.conv[0]
req_pool_indices = forward_batch.req_pool_indices req_pool_indices = forward_batch.req_pool_indices
mamba_indices = forward_batch.req_to_token_pool.get_mamba_indices( mamba_indices = get_req_to_token_pool().get_mamba_indices(req_pool_indices)
req_pool_indices
)
# Project and split into gates: B (pre-conv), C (post-conv), x (input) # Project and split into gates: B (pre-conv), C (post-conv), x (input)
proj, _ = self.in_proj(hidden_states) proj, _ = self.in_proj(hidden_states)
+3 -4
View File
@@ -42,6 +42,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
VocabParallelEmbedding, VocabParallelEmbedding,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_req_to_token_pool
from sglang.srt.model_loader.weight_utils import ( from sglang.srt.model_loader.weight_utils import (
default_weight_loader, default_weight_loader,
sharded_weight_loader, sharded_weight_loader,
@@ -326,12 +327,10 @@ class Lfm2MoeShortConv(nn.Module):
if forward_batch.forward_mode.is_idle(): if forward_batch.forward_mode.is_idle():
return hidden_states return hidden_states
layer_cache = forward_batch.req_to_token_pool.mamba2_layer_cache(self.layer_idx) layer_cache = get_req_to_token_pool().mamba2_layer_cache(self.layer_idx)
conv_state = layer_cache.conv[0] conv_state = layer_cache.conv[0]
req_pool_indices = forward_batch.req_pool_indices req_pool_indices = forward_batch.req_pool_indices
mamba_indices = forward_batch.req_to_token_pool.get_mamba_indices( mamba_indices = get_req_to_token_pool().get_mamba_indices(req_pool_indices)
req_pool_indices
)
proj, _ = self.in_proj(hidden_states) proj, _ = self.in_proj(hidden_states)
B_gate, C_gate, x = proj.chunk(3, dim=-1) B_gate, C_gate, x = proj.chunk(3, dim=-1)
+8 -4
View File
@@ -14,6 +14,10 @@ from sglang.srt.distributed import (
from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import (
get_req_to_token_pool,
get_token_to_kv_pool,
)
from sglang.srt.models.registry import import_model_classes from sglang.srt.models.registry import import_model_classes
from sglang.srt.utils import is_npu from sglang.srt.utils import is_npu
@@ -221,9 +225,9 @@ class MindSporeForCausalLM(torch.nn.Module):
def prepare_cache(cache_list, is_key_cache): def prepare_cache(cache_list, is_key_cache):
for i in range(self.config.num_hidden_layers): for i in range(self.config.num_hidden_layers):
if is_key_cache: if is_key_cache:
cache = forward_batch.token_to_kv_pool.get_key_buffer(i) cache = get_token_to_kv_pool().get_key_buffer(i)
else: else:
cache = forward_batch.token_to_kv_pool.get_value_buffer(i) cache = get_token_to_kv_pool().get_value_buffer(i)
cache_ms = tensor_torch2ms(cache) cache_ms = tensor_torch2ms(cache)
if self.use_mla and cache_ms.ndim == 3: if self.use_mla and cache_ms.ndim == 3:
cache_ms = mint.unsqueeze(cache_ms, 2) cache_ms = mint.unsqueeze(cache_ms, 2)
@@ -275,10 +279,10 @@ class MindSporeForCausalLM(torch.nn.Module):
if forward_batch.forward_mode.is_target_verify(): if forward_batch.forward_mode.is_target_verify():
q_seq_lens = q_seq_lens * forward_batch.spec_info.num_tokens_per_req q_seq_lens = q_seq_lens * forward_batch.spec_info.num_tokens_per_req
page_size = forward_batch.token_to_kv_pool.page_size page_size = get_token_to_kv_pool().page_size
block_tables = tensor_torch2ms( block_tables = tensor_torch2ms(
( (
forward_batch.req_to_token_pool.req_to_token[ get_req_to_token_pool().req_to_token[
forward_batch.req_pool_indices, : batch_valid_length.max() forward_batch.req_pool_indices, : batch_valid_length.max()
][:, ::page_size] ][:, ::page_size]
// page_size // page_size
+3 -2
View File
@@ -69,6 +69,7 @@ from sglang.srt.model_executor.breakable_cuda_graph.context import (
is_in_breakable_cuda_graph, is_in_breakable_cuda_graph,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.model_loader.weight_utils import ( from sglang.srt.model_loader.weight_utils import (
default_weight_loader, default_weight_loader,
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
@@ -414,7 +415,7 @@ class NemotronHMambaDecoderLayer(nn.Module):
) -> torch.Tensor: ) -> torch.Tensor:
"""Core Mamba forward logic, called directly or via split op.""" """Core Mamba forward logic, called directly or via split op."""
output = torch.empty_like(hidden_states) output = torch.empty_like(hidden_states)
attn_backend = forward_batch.attn_backend attn_backend = get_attn_backend()
assert isinstance(attn_backend, HybridLinearAttnBackend) assert isinstance(attn_backend, HybridLinearAttnBackend)
assert isinstance(attn_backend.linear_attn_backend, Mamba2AttnBackend) assert isinstance(attn_backend.linear_attn_backend, Mamba2AttnBackend)
attn_backend.linear_attn_backend.forward( attn_backend.linear_attn_backend.forward(
@@ -1020,7 +1021,7 @@ def nemotron_mamba2_with_output(
# In piecewise CUDA graph mode, hidden_states may be padded to the # In piecewise CUDA graph mode, hidden_states may be padded to the
# captured graph size. Slice to actual token count for Mamba forward. # captured graph size. Slice to actual token count for Mamba forward.
attn_backend = forward_batch.attn_backend attn_backend = get_attn_backend()
metadata = attn_backend.linear_attn_backend.forward_metadata metadata = attn_backend.linear_attn_backend.forward_metadata
num_actual_tokens = metadata.num_prefill_tokens + ( num_actual_tokens = metadata.num_prefill_tokens + (
metadata.num_decodes * metadata.draft_token_num metadata.num_decodes * metadata.draft_token_num
+2 -1
View File
@@ -23,6 +23,7 @@ from sglang.srt.layers.rotary_embedding.mrope import MRotaryEmbedding
from sglang.srt.layers.utils import PPMissingLayer, get_layer_id from sglang.srt.layers.utils import PPMissingLayer, get_layer_id
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
from sglang.srt.model_loader.weight_utils import ( from sglang.srt.model_loader.weight_utils import (
default_weight_loader, default_weight_loader,
maybe_remap_kv_scale_name, maybe_remap_kv_scale_name,
@@ -214,7 +215,7 @@ class Qwen3Attention(nn.Module):
qkv_3d = qkv.view(num_tokens, -1, self.head_dim) qkv_3d = qkv.view(num_tokens, -1, self.head_dim)
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = get_token_to_kv_pool()
k_cache, v_cache = token_to_kv_pool.get_kv_buffer(self.attn.layer_id) k_cache, v_cache = token_to_kv_pool.get_kv_buffer(self.attn.layer_id)
slot_mapping = forward_batch.out_cache_loc slot_mapping = forward_batch.out_cache_loc
+8 -4
View File
@@ -54,6 +54,10 @@ from sglang.srt.layers.vocab_parallel_embedding import (
) )
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_executor.forward_context import (
get_attn_backend,
get_token_to_kv_pool,
)
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.bailing_moe import BailingMoEForCausalLM from sglang.srt.models.bailing_moe import BailingMoEForCausalLM
from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha import ( from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha import (
@@ -605,7 +609,7 @@ class SarvamMoEMLAAttention(nn.Module):
self.current_attention_backend == "fa3" self.current_attention_backend == "fa3"
and self.kv_cache_dtype != "auto" and self.kv_cache_dtype != "auto"
): ):
attn_dtype = forward_batch.token_to_kv_pool.dtype attn_dtype = get_token_to_kv_pool().dtype
else: else:
attn_dtype = k_nope.dtype attn_dtype = k_nope.dtype
k = k_nope.new_empty(*k_shape, dtype=attn_dtype) k = k_nope.new_empty(*k_shape, dtype=attn_dtype)
@@ -671,7 +675,7 @@ class SarvamMoEMLAAttention(nn.Module):
q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe) q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe)
q[..., self.qk_nope_head_dim :] = q_pe q[..., self.qk_nope_head_dim :] = q_pe
forward_batch.token_to_kv_pool.set_mla_kv_buffer( get_token_to_kv_pool().set_mla_kv_buffer(
self.attn_mha, self.attn_mha,
forward_batch.out_cache_loc, forward_batch.out_cache_loc,
k_nope, k_nope,
@@ -701,8 +705,8 @@ class SarvamMoEMLAAttention(nn.Module):
forward_batch.prepare_chunked_prefix_cache_info(q.device) forward_batch.prepare_chunked_prefix_cache_info(q.device)
else: else:
forward_batch.num_prefix_chunks = 0 forward_batch.num_prefix_chunks = 0
if hasattr(forward_batch.attn_backend, "init_mha_chunk_metadata"): if hasattr(get_attn_backend(), "init_mha_chunk_metadata"):
forward_batch.attn_backend.init_mha_chunk_metadata(forward_batch) get_attn_backend().init_mha_chunk_metadata(forward_batch)
forward_batch.set_attn_attend_prefix_cache(False) forward_batch.set_attn_attend_prefix_cache(False)
forward_batch.mha_return_lse = do_prefix_merge forward_batch.mha_return_lse = do_prefix_merge
+5 -4
View File
@@ -31,6 +31,7 @@ from sglang.srt.layers.utils.cp_utils import is_prefill_context_parallel_enabled
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import get_current_device_stream_fast, is_cuda, is_hip from sglang.srt.utils import get_current_device_stream_fast, is_cuda, is_hip
@@ -275,11 +276,11 @@ class AutoWeightsLoader:
def enable_fused_set_kv_buffer(forward_batch: ForwardBatch): def enable_fused_set_kv_buffer(forward_batch: ForwardBatch):
"""Enable fused set_kv_buffer only on CUDA with bfloat16 KV cache.""" """Enable fused set_kv_buffer only on CUDA with bfloat16 KV cache."""
pool = get_token_to_kv_pool()
return ( return (
_is_cuda _is_cuda
and hasattr(forward_batch.token_to_kv_pool, "dtype") and pool.dtype == torch.bfloat16
and forward_batch.token_to_kv_pool.dtype == torch.bfloat16 and not isinstance(pool, SWAKVPool)
and not isinstance(forward_batch.token_to_kv_pool, SWAKVPool)
and not is_prefill_context_parallel_enabled() and not is_prefill_context_parallel_enabled()
) or (_is_hip and not is_prefill_context_parallel_enabled()) ) or (_is_hip and not is_prefill_context_parallel_enabled())
@@ -292,7 +293,7 @@ def create_fused_set_kv_buffer_arg(
from sglang.jit_kernel.rope import FusedSetKVBufferArg from sglang.jit_kernel.rope import FusedSetKVBufferArg
layer_id = layer.layer_id layer_id = layer.layer_id
token_to_kv_pool = forward_batch.token_to_kv_pool token_to_kv_pool = get_token_to_kv_pool()
k_buffer = token_to_kv_pool.get_key_buffer(layer_id) k_buffer = token_to_kv_pool.get_key_buffer(layer_id)
v_buffer = token_to_kv_pool.get_value_buffer(layer_id) v_buffer = token_to_kv_pool.get_value_buffer(layer_id)
@@ -646,9 +646,6 @@ class DFlashWorker:
seq_lens_sum=seq_lens_sum, seq_lens_sum=seq_lens_sum,
seq_lens_cpu=seq_lens_cpu, seq_lens_cpu=seq_lens_cpu,
positions=positions, positions=positions,
req_to_token_pool=self.draft_model_runner.req_to_token_pool,
token_to_kv_pool=self.draft_model_runner.token_to_kv_pool,
attn_backend=self.draft_model_runner.attn_backend,
input_embeds=input_embeds, input_embeds=input_embeds,
spec_algorithm=SpeculativeAlgorithm.DFLASH, spec_algorithm=SpeculativeAlgorithm.DFLASH,
spec_info=draft_spec_info, spec_info=draft_spec_info,
@@ -24,6 +24,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
) )
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
from sglang.srt.speculative.eagle_info import EagleDraftInput from sglang.srt.speculative.eagle_info import EagleDraftInput
from sglang.srt.speculative.spec_utils import ( from sglang.srt.speculative.spec_utils import (
@@ -332,8 +333,6 @@ class EAGLEDraftCudaGraphRunner:
seq_lens_cpu=seq_lens_cpu, seq_lens_cpu=seq_lens_cpu,
extend_seq_lens=extend_seq_lens, extend_seq_lens=extend_seq_lens,
extend_seq_lens_cpu=extend_seq_lens_cpu, extend_seq_lens_cpu=extend_seq_lens_cpu,
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, out_cache_loc=out_cache_loc,
seq_lens_sum=seq_lens.sum().item(), seq_lens_sum=seq_lens.sum().item(),
return_logprob=False, return_logprob=False,
@@ -350,15 +349,10 @@ class EAGLEDraftCudaGraphRunner:
), ),
) )
# Attention backend
self.draft_attn_backend.init_forward_metadata_capture_cuda_graph(forward_batch)
# Run and capture
def run_once(): def run_once():
if self.model_runner.is_hybrid_swa: if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache() self.model_runner.token_to_kv_pool.invalidate_loc_cache()
# Clean intermediate result cache for DP attention
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
set_dp_buffer_len( set_dp_buffer_len(
global_dp_buffer_len, global_dp_buffer_len,
@@ -367,7 +361,6 @@ class EAGLEDraftCudaGraphRunner:
) )
set_is_extend_in_batch(False) set_is_extend_in_batch(False)
# Backup fields that are modified in-place in `draft_forward`.
output_cache_loc_backup = forward_batch.out_cache_loc output_cache_loc_backup = forward_batch.out_cache_loc
hidden_states_backup = forward_batch.spec_info.hidden_states hidden_states_backup = forward_batch.spec_info.hidden_states
@@ -378,13 +371,15 @@ class EAGLEDraftCudaGraphRunner:
forward_batch.positions.sub_(self.eagle_worker.speculative_num_steps - 1) forward_batch.positions.sub_(self.eagle_worker.speculative_num_steps - 1)
return ret return ret
self.deepep_adapter.capture(is_extend_in_batch=False) with forward_context(ForwardContext(attn_backend=self.draft_attn_backend)):
self.draft_attn_backend.init_forward_metadata_capture_cuda_graph(
self._capture_init(run_once) forward_batch
)
out = self._capture_graph( self.deepep_adapter.capture(is_extend_in_batch=False)
graph, get_global_graph_memory_pool(), stream, run_once 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()) set_global_graph_memory_pool(graph.pool())
return graph, out return graph, out
@@ -25,6 +25,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
) )
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
from sglang.srt.speculative.spec_utils import fast_topk from sglang.srt.speculative.spec_utils import fast_topk
@@ -352,8 +353,6 @@ class EAGLEDraftExtendCudaGraphRunner:
num_accept_tokens=num_accept_tokens, num_accept_tokens=num_accept_tokens,
) )
self.deepep_adapter.capture(is_extend_in_batch=True)
# Forward batch # Forward batch
forward_batch = ForwardBatch( forward_batch = ForwardBatch(
forward_mode=self.forward_mode, forward_mode=self.forward_mode,
@@ -365,8 +364,6 @@ class EAGLEDraftExtendCudaGraphRunner:
next_token_logits_buffer=next_token_logits_buffer, next_token_logits_buffer=next_token_logits_buffer,
extend_seq_lens=extend_seq_lens, extend_seq_lens=extend_seq_lens,
extend_seq_lens_cpu=extend_seq_lens_cpu, extend_seq_lens_cpu=extend_seq_lens_cpu,
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, out_cache_loc=out_cache_loc,
seq_lens_sum=seq_lens.sum().item(), seq_lens_sum=seq_lens.sum().item(),
return_logprob=False, return_logprob=False,
@@ -379,23 +376,10 @@ class EAGLEDraftExtendCudaGraphRunner:
spec_algorithm=self.model_runner.spec_algorithm, spec_algorithm=self.model_runner.spec_algorithm,
spec_info=spec_info, spec_info=spec_info,
capture_hidden_mode=CaptureHiddenMode.LAST, capture_hidden_mode=CaptureHiddenMode.LAST,
attn_backend=self.draft_extend_attn_backend,
padded_static_len=self.padded_static_len, padded_static_len=self.padded_static_len,
) )
self.draft_extend_attn_backend.init_forward_metadata_capture_cuda_graph(
bs=bs,
num_tokens=num_tokens,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
encoder_lens=None,
forward_mode=self.forward_mode,
spec_info=spec_info,
)
# Run and capture
def run_once(): def run_once():
# model.forward() bypasses _forward_raw(), so invalidate manually.
if self.model_runner.is_hybrid_swa: if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache() self.model_runner.token_to_kv_pool.invalidate_loc_cache()
@@ -424,11 +408,23 @@ class EAGLEDraftExtendCudaGraphRunner:
forward_batch.spec_info.hidden_states = hidden_states_backup forward_batch.spec_info.hidden_states = hidden_states_backup
return ret return ret
self._capture_init(run_once) with forward_context(
ForwardContext(attn_backend=self.draft_extend_attn_backend)
out = self._capture_graph( ):
graph, get_global_graph_memory_pool(), stream, run_once self.draft_extend_attn_backend.init_forward_metadata_capture_cuda_graph(
) bs=bs,
num_tokens=num_tokens,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
encoder_lens=None,
forward_mode=self.forward_mode,
spec_info=spec_info,
)
self.deepep_adapter.capture(is_extend_in_batch=True)
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()) set_global_graph_memory_pool(graph.pool())
return graph, out return graph, out
+23 -9
View File
@@ -1,3 +1,4 @@
import contextlib
import logging import logging
import time import time
from contextlib import contextmanager from contextlib import contextmanager
@@ -30,6 +31,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
) )
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.observability.req_time_stats import set_time_batch from sglang.srt.observability.req_time_stats import set_time_batch
from sglang.srt.observability.trace import get_global_tracing_enabled from sglang.srt.observability.trace import get_global_tracing_enabled
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -881,13 +883,18 @@ class EAGLEWorker(TpModelWorker):
): ):
out_cache_loc = out_cache_loc.contiguous() out_cache_loc = out_cache_loc.contiguous()
forward_batch.out_cache_loc = out_cache_loc[i] forward_batch.out_cache_loc = out_cache_loc[i]
forward_batch.attn_backend = self.draft_attn_backend.attn_backends[i]
spec_info.hidden_states = hidden_states spec_info.hidden_states = hidden_states
# Run forward # Run forward under a per-step ForwardContext so the model layer
logits_output = self.draft_model_runner.forward( # reads attn_backends[i] for the i-th draft step. ``_forward_raw``
forward_batch, skip_attn_backend_init=True # is no-op for the attn_backend half when a context is already
).logits_output # active, so this outer wrap is what reaches RadixAttention.
with forward_context(
ForwardContext(attn_backend=self.draft_attn_backend.attn_backends[i])
):
logits_output = self.draft_model_runner.forward(
forward_batch, skip_attn_backend_init=True
).logits_output
maybe_detect_nan(logits_output.next_token_logits, f"draft_forward step {i}") maybe_detect_nan(logits_output.next_token_logits, f"draft_forward step {i}")
probs = torch.softmax(logits_output.next_token_logits, dim=-1) probs = torch.softmax(logits_output.next_token_logits, dim=-1)
topk_p, topk_index = fast_topk(probs, self.topk, dim=-1) topk_p, topk_index = fast_topk(probs, self.topk, dim=-1)
@@ -1197,16 +1204,23 @@ class EAGLEWorker(TpModelWorker):
hidden_states = logits_output.hidden_states hidden_states = logits_output.hidden_states
else: else:
forward_batch.can_run_dp_cuda_graph = False forward_batch.can_run_dp_cuda_graph = False
attn_backend = None
if not forward_batch.forward_mode.is_idle(): if not forward_batch.forward_mode.is_idle():
attn_backend = ( attn_backend = (
self.draft_extend_attn_backend self.draft_extend_attn_backend
or self.draft_model_runner.attn_backend or self.draft_model_runner.attn_backend
) )
attn_backend.init_forward_metadata(forward_batch) attn_backend.init_forward_metadata(forward_batch)
forward_batch.attn_backend = attn_backend # Publish the chosen backend via ForwardContext so model code
logits_output = self.draft_model_runner.forward( # picks it up for this forward (no runner-attr mutation).
forward_batch, skip_attn_backend_init=True if attn_backend is not None:
).logits_output ctx_mgr = forward_context(ForwardContext(attn_backend=attn_backend))
else:
ctx_mgr = contextlib.nullcontext()
with ctx_mgr:
logits_output = self.draft_model_runner.forward(
forward_batch, skip_attn_backend_init=True
).logits_output
# Non-cuda-graph path: compute topk_p / topk_index inline. # Non-cuda-graph path: compute topk_p / topk_index inline.
probs = torch.softmax(logits_output.next_token_logits, dim=-1) probs = torch.softmax(logits_output.next_token_logits, dim=-1)
topk_p, topk_index = fast_topk(probs, self.topk, dim=-1) topk_p, topk_index = fast_topk(probs, self.topk, dim=-1)
@@ -33,6 +33,7 @@ from sglang.srt.managers.scheduler import GenerationBatchResult
from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.model_executor.cuda_graph_runner import CudaGraphRunner from sglang.srt.model_executor.cuda_graph_runner import CudaGraphRunner
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode, ForwardBatch from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode, ForwardBatch
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.adaptive_runtime_state import ( from sglang.srt.speculative.adaptive_runtime_state import (
AdaptiveController, AdaptiveController,
@@ -466,13 +467,17 @@ class EagleDraftWorker(BaseDraftWorker):
# Set inputs # Set inputs
forward_batch.input_ids = input_ids forward_batch.input_ids = input_ids
forward_batch.out_cache_loc = out_cache_loc[i] forward_batch.out_cache_loc = out_cache_loc[i]
forward_batch.attn_backend = self.draft_attn_backend.attn_backends[i]
spec_info.hidden_states = hidden_states spec_info.hidden_states = hidden_states
# Run forward # Run forward under a per-step ForwardContext so the model layer
logits_output = self.draft_runner.forward( # reads attn_backends[i] for the i-th draft step. ``_forward_raw``
forward_batch, skip_attn_backend_init=True # honors the outer context and does not override.
).logits_output with forward_context(
ForwardContext(attn_backend=self.draft_attn_backend.attn_backends[i])
):
logits_output = self.draft_runner.forward(
forward_batch, skip_attn_backend_init=True
).logits_output
maybe_detect_nan(logits_output.next_token_logits, f"draft_forward step {i}") maybe_detect_nan(logits_output.next_token_logits, f"draft_forward step {i}")
probs = torch.softmax(logits_output.next_token_logits, dim=-1) probs = torch.softmax(logits_output.next_token_logits, dim=-1)
topk_p, topk_index = fast_topk(probs, self.topk, dim=-1) topk_p, topk_index = fast_topk(probs, self.topk, dim=-1)
@@ -23,6 +23,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
) )
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
from sglang.srt.speculative.frozen_kv_mtp_info import FrozenKVMTPDraftInput from sglang.srt.speculative.frozen_kv_mtp_info import FrozenKVMTPDraftInput
from sglang.srt.utils import ( from sglang.srt.utils import (
@@ -266,9 +267,6 @@ class FrozenKVMTPCudaGraphRunner:
req_pool_indices=req_pool_indices, req_pool_indices=req_pool_indices,
seq_lens=seq_lens, seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu, seq_lens_cpu=seq_lens_cpu,
req_to_token_pool=self.model_runner.req_to_token_pool,
token_to_kv_pool=self.frozen_kv_mtp_worker.kv_context.target_token_to_kv_pool,
attn_backend=self.draft_attn_backend,
out_cache_loc=None, out_cache_loc=None,
seq_lens_sum=seq_lens.sum().item(), seq_lens_sum=seq_lens.sum().item(),
return_logprob=False, return_logprob=False,
@@ -283,10 +281,6 @@ class FrozenKVMTPCudaGraphRunner:
capture_hidden_mode=CaptureHiddenMode.LAST, capture_hidden_mode=CaptureHiddenMode.LAST,
) )
self.frozen_kv_mtp_worker._init_frozen_kv_metadata_capture_cuda_graph(
forward_batch
)
def run_once(): def run_once():
if self.model_runner.is_hybrid_swa: if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache() self.model_runner.token_to_kv_pool.invalidate_loc_cache()
@@ -306,11 +300,25 @@ class FrozenKVMTPCudaGraphRunner:
forward_batch.spec_info.hidden_states = hidden_states_backup forward_batch.spec_info.hidden_states = hidden_states_backup
return ret return ret
self.deepep_adapter.capture(is_extend_in_batch=False) # Swap the draft backend's token_to_kv_pool to the frozen target pool
self._capture_init(run_once) # for the capture; the single backend-attr swap is seen by both
out = self._capture_graph( # ``get_token_to_kv_pool()`` (via ``get_attn_backend()``) and the
graph, get_global_graph_memory_pool(), stream, run_once # backend's own reads.
) target_pool = self.frozen_kv_mtp_worker.kv_context.target_token_to_kv_pool
saved_backend_pool = self.draft_attn_backend.token_to_kv_pool
self.draft_attn_backend.token_to_kv_pool = target_pool
try:
with forward_context(ForwardContext(attn_backend=self.draft_attn_backend)):
self.frozen_kv_mtp_worker._init_frozen_kv_metadata_capture_cuda_graph(
forward_batch
)
self.deepep_adapter.capture(is_extend_in_batch=False)
self._capture_init(run_once)
out = self._capture_graph(
graph, get_global_graph_memory_pool(), stream, run_once
)
finally:
self.draft_attn_backend.token_to_kv_pool = saved_backend_pool
set_global_graph_memory_pool(graph.pool()) set_global_graph_memory_pool(graph.pool())
return graph, out return graph, out
@@ -14,7 +14,7 @@
from __future__ import annotations from __future__ import annotations
from contextlib import contextmanager from contextlib import contextmanager
from typing import Tuple from typing import TYPE_CHECKING, Tuple
import torch import torch
@@ -28,39 +28,66 @@ from sglang.srt.speculative.frozen_kv_mtp_info import (
) )
from sglang.srt.speculative.spec_utils import fast_topk from sglang.srt.speculative.spec_utils import fast_topk
if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
@contextmanager @contextmanager
def frozen_kv_target_view(forward_batch: ForwardBatch, kv_context: FrozenKVMTPContext): def frozen_kv_target_view(
"""Build attention metadata against committed target-prefix geometry.""" forward_batch: ForwardBatch,
kv_context: FrozenKVMTPContext,
draft_attn_backend: "AttentionBackend",
):
"""Build attention metadata against committed target-prefix geometry.
Swaps ``draft_attn_backend.token_to_kv_pool`` to the frozen target pool
so any helper that reads ``get_token_to_kv_pool()`` during metadata init
sees the frozen target pool. Pool refs are derived from
``get_attn_backend().token_to_kv_pool`` — the single backend-attribute
swap is seen by both readers (``get_token_to_kv_pool()`` and the
backend's own ``self.token_to_kv_pool``).
"""
if kv_context is None: if kv_context is None:
raise RuntimeError( raise RuntimeError(
"Frozen-KV MTP target view called before the model was bound; " "Frozen-KV MTP target view called before the model was bound; "
"bind the frozen KV context first." "bind the frozen KV context first."
) )
saved_spec_info = forward_batch.spec_info saved_spec_info = forward_batch.spec_info
saved_kv_pool = forward_batch.token_to_kv_pool
forward_batch.spec_info = None forward_batch.spec_info = None
forward_batch.token_to_kv_pool = kv_context.target_token_to_kv_pool saved_backend_pool = draft_attn_backend.token_to_kv_pool
draft_attn_backend.token_to_kv_pool = kv_context.target_token_to_kv_pool
try: try:
yield yield
finally: finally:
forward_batch.spec_info = saved_spec_info forward_batch.spec_info = saved_spec_info
forward_batch.token_to_kv_pool = saved_kv_pool draft_attn_backend.token_to_kv_pool = saved_backend_pool
@contextmanager @contextmanager
def target_kv_pool_view(forward_batch: ForwardBatch, kv_context: FrozenKVMTPContext): def target_kv_pool_view(
forward_batch: ForwardBatch,
kv_context: FrozenKVMTPContext,
draft_attn_backend: "AttentionBackend",
):
"""Run the draft model's forward with the target's frozen KV pool.
Swaps ``draft_attn_backend.token_to_kv_pool`` to the frozen target pool.
The single backend-attribute swap is seen by both readers —
``get_token_to_kv_pool()`` (because it resolves through
``get_attn_backend()``) and the backend's own ``self.token_to_kv_pool``
reads (because ``self is draft_attn_backend``).
"""
if kv_context is None: if kv_context is None:
raise RuntimeError( raise RuntimeError(
"Frozen-KV MTP target KV pool view called before the model was bound; " "Frozen-KV MTP target KV pool view called before the model was bound; "
"bind the frozen KV context first." "bind the frozen KV context first."
) )
saved_kv_pool = forward_batch.token_to_kv_pool saved_backend_pool = draft_attn_backend.token_to_kv_pool
forward_batch.token_to_kv_pool = kv_context.target_token_to_kv_pool draft_attn_backend.token_to_kv_pool = kv_context.target_token_to_kv_pool
try: try:
yield yield
finally: finally:
forward_batch.token_to_kv_pool = saved_kv_pool draft_attn_backend.token_to_kv_pool = saved_backend_pool
def set_frozen_kv_positions(forward_batch: ForwardBatch, topk: int) -> None: def set_frozen_kv_positions(forward_batch: ForwardBatch, topk: int) -> None:
@@ -39,6 +39,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
) )
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
from sglang.srt.observability.req_time_stats import set_time_batch from sglang.srt.observability.req_time_stats import set_time_batch
from sglang.srt.observability.trace import get_global_tracing_enabled from sglang.srt.observability.trace import get_global_tracing_enabled
@@ -248,10 +249,14 @@ class FrozenKVMTPWorker(TpModelWorker):
self.kv_context = ctx self.kv_context = ctx
def _frozen_kv_target_view(self, forward_batch: ForwardBatch): def _frozen_kv_target_view(self, forward_batch: ForwardBatch):
return frozen_kv_target_view(forward_batch, self.kv_context) return frozen_kv_target_view(
forward_batch, self.kv_context, self.draft_attn_backend
)
def _target_kv_pool_view(self, forward_batch: ForwardBatch): def _target_kv_pool_view(self, forward_batch: ForwardBatch):
return target_kv_pool_view(forward_batch, self.kv_context) return target_kv_pool_view(
forward_batch, self.kv_context, self.draft_attn_backend
)
def _set_positions(self, forward_batch: ForwardBatch) -> None: def _set_positions(self, forward_batch: ForwardBatch) -> None:
set_frozen_kv_positions(forward_batch, self.topk) set_frozen_kv_positions(forward_batch, self.topk)
@@ -275,7 +280,6 @@ class FrozenKVMTPWorker(TpModelWorker):
forward_batch.seq_lens_sum = torch.sum(forward_batch.seq_lens).item() forward_batch.seq_lens_sum = torch.sum(forward_batch.seq_lens).item()
with self._frozen_kv_target_view(forward_batch): with self._frozen_kv_target_view(forward_batch):
self.draft_attn_backend.init_forward_metadata(forward_batch) self.draft_attn_backend.init_forward_metadata(forward_batch)
forward_batch.attn_backend = self.draft_attn_backend
def _init_frozen_kv_metadata_capture_cuda_graph( def _init_frozen_kv_metadata_capture_cuda_graph(
self, forward_batch: ForwardBatch self, forward_batch: ForwardBatch
@@ -290,7 +294,6 @@ class FrozenKVMTPWorker(TpModelWorker):
forward_mode=ForwardMode.DECODE, forward_mode=ForwardMode.DECODE,
spec_info=None, spec_info=None,
) )
forward_batch.attn_backend = self.draft_attn_backend
def _init_frozen_kv_metadata_replay_cuda_graph( def _init_frozen_kv_metadata_replay_cuda_graph(
self, forward_batch: ForwardBatch, bs: int, seq_lens_sum: int self, forward_batch: ForwardBatch, bs: int, seq_lens_sum: int
@@ -310,7 +313,6 @@ class FrozenKVMTPWorker(TpModelWorker):
else None else None
), ),
) )
forward_batch.attn_backend = self.draft_attn_backend
def init_cuda_graphs(self) -> None: def init_cuda_graphs(self) -> None:
if self.server_args.disable_cuda_graph or self.speculative_num_steps <= 1: if self.server_args.disable_cuda_graph or self.speculative_num_steps <= 1:
@@ -396,7 +398,9 @@ class FrozenKVMTPWorker(TpModelWorker):
forward_batch.mm_input_embeds = mm_input_embeds forward_batch.mm_input_embeds = mm_input_embeds
self._set_positions(forward_batch) self._set_positions(forward_batch)
self._init_frozen_kv_metadata(forward_batch) self._init_frozen_kv_metadata(forward_batch)
with self._target_kv_pool_view(forward_batch): with self._target_kv_pool_view(forward_batch), forward_context(
ForwardContext(attn_backend=self.draft_attn_backend)
):
logits_output = self.draft_model_runner.forward( logits_output = self.draft_model_runner.forward(
forward_batch, skip_attn_backend_init=True forward_batch, skip_attn_backend_init=True
).logits_output ).logits_output
@@ -678,7 +682,9 @@ class FrozenKVMTPWorker(TpModelWorker):
forward_batch.spec_info.hidden_states = hidden_states forward_batch.spec_info.hidden_states = hidden_states
self._set_positions(forward_batch) self._set_positions(forward_batch)
with self._target_kv_pool_view(forward_batch): with self._target_kv_pool_view(forward_batch), forward_context(
ForwardContext(attn_backend=self.draft_attn_backend)
):
logits_output = self.draft_model_runner.forward( logits_output = self.draft_model_runner.forward(
forward_batch, skip_attn_backend_init=True forward_batch, skip_attn_backend_init=True
).logits_output ).logits_output
@@ -40,6 +40,11 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode, ForwardMode,
) )
from sglang.srt.model_executor.forward_context import (
ForwardContext,
forward_context,
get_req_to_token_pool,
)
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
from sglang.srt.speculative.multi_layer_eagle_utils import assign_new_state_triton from sglang.srt.speculative.multi_layer_eagle_utils import assign_new_state_triton
@@ -369,8 +374,6 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
seq_lens=seq_lens, seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu, seq_lens_cpu=seq_lens_cpu,
next_token_logits_buffer=next_token_logits_buffer, 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, out_cache_loc=out_cache_loc,
seq_lens_sum=seq_lens.sum().item(), seq_lens_sum=seq_lens.sum().item(),
return_logprob=False, return_logprob=False,
@@ -383,7 +386,6 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
spec_algorithm=self.model_runner.spec_algorithm, spec_algorithm=self.model_runner.spec_algorithm,
spec_info=spec_info, spec_info=spec_info,
capture_hidden_mode=capture_mode, capture_hidden_mode=capture_mode,
attn_backend=self.eagle_worker.draft_extend_attn_backend_list[self.step],
extend_seq_lens=extend_seq_lens, extend_seq_lens=extend_seq_lens,
extend_seq_lens_cpu=extend_seq_lens_cpu, extend_seq_lens_cpu=extend_seq_lens_cpu,
padded_static_len=self.padded_static_len, padded_static_len=self.padded_static_len,
@@ -400,26 +402,11 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
graph = self._create_graph() graph = self._create_graph()
stream = self.stream stream = self.stream
self.deepep_adapter.capture(is_extend_in_batch=True)
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_bs
forward_batch = self.get_forward_batch(bs) forward_batch = self.get_forward_batch(bs)
attn_backend = self.eagle_worker.draft_extend_attn_backend_list[self.step]
self.eagle_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(): def run_once():
# model.forward() bypasses _forward_raw(), so invalidate manually.
if self.model_runner.is_hybrid_swa: if self.model_runner.is_hybrid_swa:
self.model_runner.token_to_kv_pool.invalidate_loc_cache() self.model_runner.token_to_kv_pool.invalidate_loc_cache()
@@ -490,18 +477,28 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
forward_batch.batch_size, forward_batch.batch_size,
self.step, self.step,
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.req_to_token_pool.req_to_token, get_req_to_token_pool().req_to_token,
self.eagle_worker.req_to_hidden_states_pool, self.eagle_worker.req_to_hidden_states_pool,
) )
forward_batch.out_cache_loc = output_cache_loc_backup forward_batch.out_cache_loc = output_cache_loc_backup
forward_batch.spec_info.hidden_states = hidden_states_backup forward_batch.spec_info.hidden_states = hidden_states_backup
return ret return ret
self._capture_init(run_once) with forward_context(ForwardContext(attn_backend=attn_backend)):
attn_backend.init_forward_metadata_capture_cuda_graph(
out = self._capture_graph( bs=bs,
graph, get_global_graph_memory_pool(), stream, run_once 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,
)
self.deepep_adapter.capture(is_extend_in_batch=True)
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()) set_global_graph_memory_pool(graph.pool())
return graph, out return graph, out
@@ -419,9 +419,6 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
topk_p_list = [] topk_p_list = []
topk_index_list = [] topk_index_list = []
for step in range(self.speculative_num_steps): for step in range(self.speculative_num_steps):
forward_batch.req_to_token_pool = self.draft_runner_list[
step
].req_to_token_pool
output: ModelRunnerOutput = self.draft_runner_list[step].forward( output: ModelRunnerOutput = self.draft_runner_list[step].forward(
forward_batch forward_batch
) )
@@ -526,9 +523,6 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
draft_logits_output.topk_index, draft_logits_output.topk_index,
) )
else: else:
forward_batch.req_to_token_pool = self.draft_runner_list[
step
].req_to_token_pool
draft_logits_output = self.draft_runner_list[step].forward( draft_logits_output = self.draft_runner_list[step].forward(
forward_batch, skip_attn_backend_init=True forward_batch, skip_attn_backend_init=True
) )
@@ -11,6 +11,10 @@ from sglang.srt.layers.attention.torch_native_backend import TorchNativeAttnBack
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.forward_context import (
ForwardContext,
set_forward_context,
)
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -109,6 +113,9 @@ class TestFlashAttentionBackend(CustomTestCase):
self.backend = FlashAttentionBackend(self.model_runner) self.backend = FlashAttentionBackend(self.model_runner)
self.ref_backend = TorchNativeAttnBackend(self.model_runner) self.ref_backend = TorchNativeAttnBackend(self.model_runner)
self.model_runner.model_config.num_attention_heads = self.num_heads self.model_runner.model_config.num_attention_heads = self.num_heads
# Publish the backend for any RadixAttention.forward path the tests
# exercise; tearDown is unnecessary here since each test re-inits.
set_forward_context(ForwardContext(attn_backend=self.backend))
def _mock_write_to_req_to_token_pool(self, batch_size, seq_len, page_size): def _mock_write_to_req_to_token_pool(self, batch_size, seq_len, page_size):
# if page_size > 1, the token pool stores the index to the page. # if page_size > 1, the token pool stores the index to the page.
@@ -223,7 +230,6 @@ class TestFlashAttentionBackend(CustomTestCase):
extend_seq_lens_cpu=torch.tensor( extend_seq_lens_cpu=torch.tensor(
[q_len] * self.batch_size, device="cpu" [q_len] * self.batch_size, device="cpu"
), ),
attn_backend=self.backend,
) )
if attn_cp_size > 1: if attn_cp_size > 1:
forward_batch.attn_cp_metadata = type( forward_batch.attn_cp_metadata = type(
@@ -273,16 +279,11 @@ class TestFlashAttentionBackend(CustomTestCase):
[total_len] * self.batch_size, device=self.device [total_len] * self.batch_size, device=self.device
), ),
seq_lens_cpu=torch.tensor([total_len] * self.batch_size, device="cpu"), seq_lens_cpu=torch.tensor([total_len] * self.batch_size, device="cpu"),
attn_backend=self.backend,
) )
# Add token pool # Pool refs are resolved via the active ForwardContext (published in
forward_batch.req_to_token_pool = self.model_runner.req_to_token_pool # setUp). Write the test fixture's req_to_token mapping.
# Write current batch's req_to_token to req_to_token_pool
self._mock_write_to_req_to_token_pool(self.batch_size, total_len, page_size) self._mock_write_to_req_to_token_pool(self.batch_size, total_len, page_size)
# Add kv pool for this forward batch
forward_batch.token_to_kv_pool = self.model_runner.token_to_kv_pool
return forward_batch return forward_batch
@@ -307,7 +308,7 @@ class TestFlashAttentionBackend(CustomTestCase):
) )
# Set the prefix KV cache # Set the prefix KV cache
forward_batch.token_to_kv_pool.set_kv_buffer( self.model_runner.token_to_kv_pool.set_kv_buffer(
layer, layer,
torch.arange(self.batch_size * cache_len, device=self.device), torch.arange(self.batch_size * cache_len, device=self.device),
cache_k, cache_k,
@@ -8,6 +8,10 @@ from sglang.srt.layers.attention.torch_native_backend import TorchNativeAttnBack
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.forward_context import (
ForwardContext,
set_forward_context,
)
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -112,6 +116,8 @@ class TestFlashAttentionMLABackend(CustomTestCase):
self.backend = FlashAttentionBackend(self.model_runner) self.backend = FlashAttentionBackend(self.model_runner)
self.ref_backend = TorchNativeAttnBackend(self.model_runner) self.ref_backend = TorchNativeAttnBackend(self.model_runner)
self.num_local_heads = 2 self.num_local_heads = 2
# Publish the backend so RadixAttention.forward resolves correctly.
set_forward_context(ForwardContext(attn_backend=self.backend))
def _init_model_runner(self): def _init_model_runner(self):
self.model_runner = MockModelRunner( self.model_runner = MockModelRunner(
@@ -192,7 +198,6 @@ class TestFlashAttentionMLABackend(CustomTestCase):
extend_seq_lens_cpu=torch.tensor( extend_seq_lens_cpu=torch.tensor(
[q_len] * self.batch_size, device="cpu" [q_len] * self.batch_size, device="cpu"
), ),
attn_backend=self.backend,
) )
else: # ForwardMode.DECODE else: # ForwardMode.DECODE
@@ -216,15 +221,10 @@ class TestFlashAttentionMLABackend(CustomTestCase):
[total_len] * self.batch_size, device=self.device [total_len] * self.batch_size, device=self.device
), ),
seq_lens_cpu=torch.tensor([total_len] * self.batch_size, device="cpu"), seq_lens_cpu=torch.tensor([total_len] * self.batch_size, device="cpu"),
attn_backend=self.backend,
) )
# Add token pool from model runner to forward batch # Pool refs are resolved via the active ForwardContext (published in
forward_batch.req_to_token_pool = self.model_runner.req_to_token_pool # setUp); the fixture no longer needs to attach them to forward_batch.
# Add KV cache from model runner to forward batch
forward_batch.token_to_kv_pool = self.model_runner.token_to_kv_pool
return forward_batch return forward_batch
def _setup_kv_cache(self, forward_batch, layer, cache_len): def _setup_kv_cache(self, forward_batch, layer, cache_len):
@@ -250,7 +250,7 @@ class TestFlashAttentionMLABackend(CustomTestCase):
) )
# Set the prefix KV cache using MLA-specific method # Set the prefix KV cache using MLA-specific method
forward_batch.token_to_kv_pool.set_mla_kv_buffer( self.model_runner.token_to_kv_pool.set_mla_kv_buffer(
layer, layer,
torch.arange(self.batch_size * cache_len, device=self.device), torch.arange(self.batch_size * cache_len, device=self.device),
cache_k_nope, cache_k_nope,
@@ -110,12 +110,10 @@ class MockReqToTokenPool:
# Test correctness of triton kernel for computing kv indices # Test correctness of triton kernel for computing kv indices
def check_kv_indices(forward_batch): def check_kv_indices(forward_batch, req_to_token_pool):
for i in range(forward_batch.num_prefix_chunks): for i in range(forward_batch.num_prefix_chunks):
computed_kv_indices = forward_batch.prefix_chunk_kv_indices[i] computed_kv_indices = forward_batch.prefix_chunk_kv_indices[i]
req_to_token = forward_batch.req_to_token_pool.req_to_token[ req_to_token = req_to_token_pool.req_to_token[: forward_batch.batch_size, :]
: forward_batch.batch_size, :
]
ref_kv_indices = torch.empty( ref_kv_indices = torch.empty(
forward_batch.prefix_chunk_num_tokens[i], forward_batch.prefix_chunk_num_tokens[i],
dtype=torch.int32, dtype=torch.int32,
@@ -205,8 +203,20 @@ class TestPrefixChunkInfo(CustomTestCase):
extend_prefix_lens=prefix_lens, extend_prefix_lens=prefix_lens,
extend_prefix_lens_cpu=prefix_lens_cpu, extend_prefix_lens_cpu=prefix_lens_cpu,
) )
forward_batch.req_to_token_pool = self.req_to_token_pool # Pool refs are resolved via the active ForwardContext; mock an
forward_batch.token_to_kv_pool = self.token_to_kv_pool # attn_backend that carries the pools (Pattern A invariant).
from types import SimpleNamespace
from sglang.srt.model_executor.forward_context import (
ForwardContext,
set_forward_context,
)
mock_backend = SimpleNamespace(
req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool=self.token_to_kv_pool,
)
set_forward_context(ForwardContext(attn_backend=mock_backend))
forward_batch.prepare_chunked_prefix_cache_info(self.device) forward_batch.prepare_chunked_prefix_cache_info(self.device)
assert forward_batch.get_max_chunk_capacity() == max_chunk_capacity assert forward_batch.get_max_chunk_capacity() == max_chunk_capacity
@@ -221,7 +231,7 @@ class TestPrefixChunkInfo(CustomTestCase):
test_case["prefix_chunk_seq_lens"].to(self.device), test_case["prefix_chunk_seq_lens"].to(self.device),
) )
check_kv_indices(forward_batch) check_kv_indices(forward_batch, self.req_to_token_pool)
if __name__ == "__main__": if __name__ == "__main__":
@@ -20,6 +20,10 @@ from sglang.srt.layers.attention.utils import get_num_page_per_block_flashmla
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.forward_context import (
ForwardContext,
set_forward_context,
)
from sglang.srt.server_args import ( from sglang.srt.server_args import (
ServerArgs, ServerArgs,
get_global_server_args, get_global_server_args,
@@ -434,10 +438,9 @@ class TestTRTLLMMLA(CustomTestCase):
req_pool_indices=torch.arange(batch_size, device=config["device"]), req_pool_indices=torch.arange(batch_size, device=config["device"]),
seq_lens=seq_lens, seq_lens=seq_lens,
seq_lens_cpu=seq_lens.cpu(), seq_lens_cpu=seq_lens.cpu(),
attn_backend=backend,
) )
fb.req_to_token_pool = model_runner.req_to_token_pool # Publish backend for RadixAttention dispatch.
fb.token_to_kv_pool = model_runner.token_to_kv_pool set_forward_context(ForwardContext(attn_backend=backend))
# Add position information for RoPE # Add position information for RoPE
fb.positions = torch.arange(batch_size, device=config["device"]) fb.positions = torch.arange(batch_size, device=config["device"])
@@ -1167,10 +1170,9 @@ class TestTRTLLMMLA(CustomTestCase):
seq_lens_cpu=seq_lens.cpu(), seq_lens_cpu=seq_lens.cpu(),
attn_attend_prefix_cache=False, attn_attend_prefix_cache=False,
mha_return_lse=False, mha_return_lse=False,
attn_backend=backend,
) )
fb.req_to_token_pool = model_runner.req_to_token_pool # Publish backend for RadixAttention dispatch.
fb.token_to_kv_pool = model_runner.token_to_kv_pool set_forward_context(ForwardContext(attn_backend=backend))
# Add position information for RoPE # Add position information for RoPE
fb.positions = torch.arange(batch_size, device=config["device"]) fb.positions = torch.arange(batch_size, device=config["device"])
+10 -5
View File
@@ -360,7 +360,6 @@ class TestDSAIndexer(CustomTestCase):
), ),
extend_seq_lens=torch.tensor([q_len] * batch_size, device=self.device), extend_seq_lens=torch.tensor([q_len] * batch_size, device=self.device),
extend_seq_lens_cpu=torch.tensor([q_len] * batch_size, device="cpu"), extend_seq_lens_cpu=torch.tensor([q_len] * batch_size, device="cpu"),
attn_backend=self.backend,
) )
else: # ForwardMode.DECODE else: # ForwardMode.DECODE
decode_len = 1 decode_len = 1
@@ -379,12 +378,18 @@ class TestDSAIndexer(CustomTestCase):
req_pool_indices=torch.arange(batch_size, device=self.device), req_pool_indices=torch.arange(batch_size, device=self.device),
seq_lens=torch.tensor([total_len] * batch_size, device=self.device), seq_lens=torch.tensor([total_len] * batch_size, device=self.device),
seq_lens_cpu=torch.tensor([total_len] * batch_size, device="cpu"), seq_lens_cpu=torch.tensor([total_len] * batch_size, device="cpu"),
attn_backend=self.backend,
) )
# Add token pools # Pool refs + attn_backend are now resolved via the ForwardContext;
forward_batch.req_to_token_pool = self.model_runner.req_to_token_pool # publish ``self.backend`` for the duration of this fixture call so
forward_batch.token_to_kv_pool = self.model_runner.token_to_kv_pool # ``get_attn_backend()`` / ``get_token_to_kv_pool()`` /
# ``get_req_to_token_pool()`` resolve correctly.
from sglang.srt.model_executor.forward_context import (
ForwardContext,
set_forward_context,
)
set_forward_context(ForwardContext(attn_backend=self.backend))
# Mock write to req_to_token_pool # Mock write to req_to_token_pool
page_size = self.model_runner.page_size page_size = self.model_runner.page_size