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:
co-authored by
Claude Sonnet 4.6
parent
44ec2ee18d
commit
c5251a98a9
@@ -1,16 +1,31 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, Generator, List, Sequence, Union
|
||||
from contextlib import contextmanager, nullcontext
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
Dict,
|
||||
Generator,
|
||||
List,
|
||||
Optional,
|
||||
Sequence,
|
||||
Union,
|
||||
)
|
||||
|
||||
import torch
|
||||
|
||||
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:
|
||||
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")))
|
||||
|
||||
@@ -39,10 +54,15 @@ def execute_overlapped_operations(
|
||||
assert delta_stage_a == 0
|
||||
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_b = _convert_operations_to_stages(operations_b)
|
||||
executor_a = _StageExecutor("a", stages_a, inputs=inputs_a)
|
||||
executor_b = _StageExecutor("b", stages_b, inputs=inputs_b)
|
||||
executor_a = _StageExecutor("a", stages_a, inputs=inputs_a, child_ctx=child_ctx_a)
|
||||
executor_b = _StageExecutor("b", stages_b, inputs=inputs_b, child_ctx=child_ctx_b)
|
||||
|
||||
for _ in range(delta_stage):
|
||||
executor_a.next()
|
||||
@@ -58,6 +78,25 @@ def execute_overlapped_operations(
|
||||
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:
|
||||
pass
|
||||
|
||||
@@ -73,12 +112,23 @@ Stage = List[ExecutionOperation]
|
||||
|
||||
|
||||
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._stages = stages
|
||||
self._index = 0
|
||||
self._stage_state = _StateDict()
|
||||
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
|
||||
forward_batch: ForwardBatch = inputs["forward_batch"]
|
||||
@@ -102,7 +152,12 @@ class _StageExecutor:
|
||||
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:
|
||||
with _annotate_region(debug_name=op.debug_name):
|
||||
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.layers import deep_gemm_wrapper
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.communicator import (
|
||||
CommunicateContext,
|
||||
CommunicateSummableTensorPairFn,
|
||||
@@ -40,6 +39,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardMode,
|
||||
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.speculative.spec_info import SpecInput
|
||||
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}"
|
||||
)
|
||||
|
||||
assert isinstance(batch.attn_backend, TboAttnBackend)
|
||||
attn_backend_child_a, attn_backend_child_b = batch.attn_backend.children
|
||||
# Sanity check: the global attn_backend should be a TboAttnBackend
|
||||
# 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] = (
|
||||
tbo_children_num_token_non_padded
|
||||
@@ -525,7 +527,6 @@ class TboForwardBatchPreparer:
|
||||
if is_enable_two_chunk
|
||||
else batch.tbo_split_seq_index
|
||||
),
|
||||
output_attn_backend=attn_backend_child_a,
|
||||
out_num_token_non_padded=out_num_token_non_padded_a,
|
||||
)
|
||||
child_b = cls.filter_batch(
|
||||
@@ -534,7 +535,6 @@ class TboForwardBatchPreparer:
|
||||
end_token_index=batch.input_ids.shape[0],
|
||||
start_seq_index=batch.tbo_split_seq_index,
|
||||
end_seq_index=batch.batch_size,
|
||||
output_attn_backend=attn_backend_child_b,
|
||||
out_num_token_non_padded=out_num_token_non_padded_b,
|
||||
)
|
||||
|
||||
@@ -620,7 +620,6 @@ class TboForwardBatchPreparer:
|
||||
end_token_index: int,
|
||||
start_seq_index: int,
|
||||
end_seq_index: int,
|
||||
output_attn_backend: AttentionBackend,
|
||||
out_num_token_non_padded: torch.Tensor,
|
||||
):
|
||||
assert (
|
||||
@@ -692,8 +691,6 @@ class TboForwardBatchPreparer:
|
||||
"is_extend_in_batch",
|
||||
"all_extend_in_batch",
|
||||
"return_logprob",
|
||||
"req_to_token_pool",
|
||||
"token_to_kv_pool",
|
||||
"can_run_dp_cuda_graph",
|
||||
"dp_padding_mode",
|
||||
"global_forward_mode",
|
||||
@@ -743,7 +740,6 @@ class TboForwardBatchPreparer:
|
||||
else None
|
||||
),
|
||||
extend_num_tokens=extend_num_tokens,
|
||||
attn_backend=output_attn_backend,
|
||||
num_token_non_padded=out_num_token_non_padded,
|
||||
# TODO: handle it when we need TBO + DeepSeek V3.2
|
||||
num_token_non_padded_cpu=None,
|
||||
|
||||
@@ -264,11 +264,11 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
|
||||
else forward_batch.encoder_out_cache_loc
|
||||
)
|
||||
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
|
||||
)
|
||||
else:
|
||||
forward_batch.token_to_kv_pool.set_mla_kv_buffer(
|
||||
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
k,
|
||||
@@ -357,9 +357,7 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
|
||||
can_run_tbo=forward_batch.can_run_tbo,
|
||||
)
|
||||
if not self.use_mla:
|
||||
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
|
||||
layer.layer_id
|
||||
)
|
||||
key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
||||
|
||||
key_cache = key_cache.view(
|
||||
-1, self.page_size, layer.tp_k_head_num, layer.head_dim
|
||||
@@ -555,9 +553,9 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
|
||||
return output, lse
|
||||
return output
|
||||
else:
|
||||
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer(
|
||||
layer.layer_id
|
||||
).to(q.dtype)
|
||||
kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
|
||||
q.dtype
|
||||
)
|
||||
k_rope = kv_cache[:, :, layer.v_head_dim :]
|
||||
c_kv = kv_cache[:, :, : layer.v_head_dim]
|
||||
k_rope_cache = k_rope.view(
|
||||
@@ -657,11 +655,11 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
|
||||
else forward_batch.encoder_out_cache_loc
|
||||
)
|
||||
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
|
||||
)
|
||||
else:
|
||||
forward_batch.token_to_kv_pool.set_mla_kv_buffer(
|
||||
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
k,
|
||||
@@ -710,9 +708,7 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
|
||||
can_run_tbo=forward_batch.can_run_tbo,
|
||||
)
|
||||
if not self.use_mla:
|
||||
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
|
||||
layer.layer_id
|
||||
)
|
||||
key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
||||
key_cache = key_cache.view(
|
||||
-1, self.page_size, layer.tp_k_head_num, layer.head_dim
|
||||
)
|
||||
@@ -831,9 +827,7 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
|
||||
else:
|
||||
o = result
|
||||
else:
|
||||
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
|
||||
q.dtype
|
||||
)
|
||||
kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(q.dtype)
|
||||
k_rope = kv_cache[:, :, layer.v_head_dim :]
|
||||
c_kv = kv_cache[:, :, : layer.v_head_dim]
|
||||
k_rope_cache = k_rope.view(
|
||||
|
||||
@@ -208,7 +208,9 @@ class AscendAttnMaskBuilder:
|
||||
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.
|
||||
|
||||
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)
|
||||
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,
|
||||
cache_loc,
|
||||
key_cache_full,
|
||||
@@ -287,6 +289,10 @@ class AscendAttnBackend(AttentionBackend):
|
||||
self.native_attn = AscendTorchNativeAttnBackend()
|
||||
self.graph_metadata = {}
|
||||
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.graph_mode = 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
|
||||
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
|
||||
][:, :: self.page_size]
|
||||
// self.page_size
|
||||
@@ -366,7 +372,7 @@ class AscendAttnBackend(AttentionBackend):
|
||||
self.forward_metadata.block_tables_swa = (
|
||||
(
|
||||
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
|
||||
]
|
||||
][:, :: self.page_size]
|
||||
@@ -421,7 +427,7 @@ class AscendAttnBackend(AttentionBackend):
|
||||
for req_idx, seq_len in zip(
|
||||
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_indices[:seq_len][:: self.page_size] // self.page_size
|
||||
)
|
||||
@@ -883,11 +889,11 @@ class AscendAttnBackend(AttentionBackend):
|
||||
if save_kv_cache:
|
||||
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)
|
||||
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
|
||||
)
|
||||
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 self.forward_metadata.actual_seq_lengths_q is not None:
|
||||
@@ -1041,7 +1047,12 @@ class AscendAttnBackend(AttentionBackend):
|
||||
if is_cp_mode:
|
||||
# All-gather K/V from all CP ranks and write full sequence to KV pool
|
||||
_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:
|
||||
# support cross attention
|
||||
@@ -1050,10 +1061,10 @@ class AscendAttnBackend(AttentionBackend):
|
||||
if not layer.is_cross_attention
|
||||
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)
|
||||
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
|
||||
if sinks is not None:
|
||||
# Use SWA block tables if hybrid SWA is enabled for this layer
|
||||
@@ -1200,7 +1211,7 @@ class AscendAttnBackend(AttentionBackend):
|
||||
o_,
|
||||
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),
|
||||
forward_batch.req_to_token_pool.req_to_token,
|
||||
self.req_to_token_pool.req_to_token,
|
||||
forward_batch.req_pool_indices,
|
||||
forward_batch.seq_lens,
|
||||
forward_batch.extend_prefix_lens,
|
||||
@@ -1223,10 +1234,8 @@ class AscendAttnBackend(AttentionBackend):
|
||||
if layer.qk_head_dim == layer.v_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)
|
||||
v_buffer = forward_batch.token_to_kv_pool.get_value_buffer(
|
||||
layer.layer_id
|
||||
)
|
||||
k_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_buffer = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
kv_cached = torch.index_select(
|
||||
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
|
||||
k_buffer = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_buffer = forward_batch.token_to_kv_pool.get_value_buffer(
|
||||
layer.layer_id
|
||||
)
|
||||
k_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_buffer = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
kv_cached = torch.index_select(
|
||||
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_c, k_rope = k.split([kv_lora_rank, self.qk_rope_head_dim], dim=-1)
|
||||
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
|
||||
)
|
||||
attn_output = q.new_empty(
|
||||
@@ -1435,17 +1442,15 @@ class AscendAttnBackend(AttentionBackend):
|
||||
)
|
||||
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)
|
||||
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(
|
||||
layer.layer_id
|
||||
)
|
||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
kv_cache = torch.cat([k_cache, v_cache], dim=-1)
|
||||
attn_output = self.native_attn.run_sdpa_forward_extend(
|
||||
q,
|
||||
attn_output,
|
||||
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),
|
||||
forward_batch.req_to_token_pool.req_to_token,
|
||||
self.req_to_token_pool.req_to_token,
|
||||
forward_batch.req_pool_indices,
|
||||
forward_batch.seq_lens,
|
||||
forward_batch.extend_prefix_lens,
|
||||
@@ -1514,12 +1519,12 @@ class AscendAttnBackend(AttentionBackend):
|
||||
topk_indices: Optional[torch.Tensor] = None,
|
||||
):
|
||||
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
|
||||
)
|
||||
|
||||
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
k_cache = self.token_to_kv_pool.get_key_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)
|
||||
|
||||
if self.forward_metadata.seq_lens_cpu_int is None:
|
||||
@@ -1574,21 +1579,21 @@ class AscendAttnBackend(AttentionBackend):
|
||||
if self.use_mla:
|
||||
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)
|
||||
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
|
||||
)
|
||||
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
|
||||
)
|
||||
|
||||
if not self.use_mla:
|
||||
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(
|
||||
layer.layer_id
|
||||
).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(
|
||||
layer.layer_id
|
||||
).view(-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim)
|
||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).view(
|
||||
-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
|
||||
)
|
||||
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id).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()
|
||||
if not self.graph_mode:
|
||||
num_token_padding = query.shape[0]
|
||||
@@ -1642,7 +1647,7 @@ class AscendAttnBackend(AttentionBackend):
|
||||
)
|
||||
return attn_output
|
||||
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():
|
||||
k_rope_cache = _reshape_kv_for_fia_nz(
|
||||
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:
|
||||
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)
|
||||
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
|
||||
)
|
||||
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
|
||||
)
|
||||
|
||||
if sinks is not None:
|
||||
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
k_cache = self.token_to_kv_pool.get_key_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
|
||||
if self.is_hybrid_swa and layer.sliding_window_size != -1:
|
||||
@@ -1788,12 +1793,12 @@ class AscendAttnBackend(AttentionBackend):
|
||||
return attn_out
|
||||
|
||||
if not self.use_mla:
|
||||
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(
|
||||
layer.layer_id
|
||||
).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(
|
||||
layer.layer_id
|
||||
).view(-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim)
|
||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).view(
|
||||
-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
|
||||
)
|
||||
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id).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)
|
||||
if self.forward_metadata.seq_lens_cpu_int is None:
|
||||
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)
|
||||
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():
|
||||
k_rope_cache = _reshape_kv_for_fia_nz(
|
||||
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
|
||||
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]
|
||||
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
|
||||
if sinks is not None:
|
||||
# Use SWA block tables if hybrid SWA is enabled for this layer
|
||||
@@ -2098,7 +2103,7 @@ class AscendAttnBackend(AttentionBackend):
|
||||
o_,
|
||||
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),
|
||||
forward_batch.req_to_token_pool.req_to_token,
|
||||
self.req_to_token_pool.req_to_token,
|
||||
forward_batch.req_pool_indices,
|
||||
forward_batch.seq_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)
|
||||
else:
|
||||
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
|
||||
)
|
||||
num_tokens = q.shape[0]
|
||||
kv_c = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
k_pe = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
kv_c = self.token_to_kv_pool.get_key_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:
|
||||
"""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."
|
||||
)
|
||||
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
|
||||
)
|
||||
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
v_cache = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
num_block, block_size, _, _ = k_cache.shape
|
||||
key = k_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
|
||||
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -253,7 +257,7 @@ class NPUFusedMLAPreprocess(torch.nn.Module):
|
||||
return cos, sin
|
||||
|
||||
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)
|
||||
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"
|
||||
self.kvCache = self.kvCache.view(
|
||||
-1,
|
||||
forward_batch.attn_backend.page_size,
|
||||
get_attn_backend().page_size,
|
||||
1,
|
||||
forward_batch.attn_backend.kv_lora_rank,
|
||||
get_attn_backend().kv_lora_rank,
|
||||
)
|
||||
self.kvCacheRope = self.kvCacheRope.view(
|
||||
-1,
|
||||
forward_batch.attn_backend.page_size,
|
||||
get_attn_backend().page_size,
|
||||
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(
|
||||
latent_cache,
|
||||
|
||||
@@ -16,6 +16,7 @@ from sglang.srt.layers.attention.dsa.utils import (
|
||||
dsa_use_prefill_cp,
|
||||
)
|
||||
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:
|
||||
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)
|
||||
|
||||
ckv_cache, k_rope_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
|
||||
m.layer_id
|
||||
)
|
||||
ckv_cache, k_rope_cache = get_token_to_kv_pool().get_kv_buffer(m.layer_id)
|
||||
_, _, 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
|
||||
m.kv_a_layernorm.weight,
|
||||
@@ -115,7 +114,7 @@ def forward_mha_prepare_npu(
|
||||
if m.rotary_emb is not None:
|
||||
q_pe, k_pe = m.rotary_emb(positions, q_pe, k_pe)
|
||||
# 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
|
||||
)
|
||||
|
||||
|
||||
@@ -204,6 +204,11 @@ class AiterAttnBackend(AttentionBackend):
|
||||
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
|
||||
self.use_sliding_window_kv_pool = (
|
||||
isinstance(model_runner.token_to_kv_pool, SWAKVPool)
|
||||
@@ -211,7 +216,6 @@ class AiterAttnBackend(AttentionBackend):
|
||||
)
|
||||
|
||||
if self.use_sliding_window_kv_pool:
|
||||
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||
self.use_triton_unified_attention = True
|
||||
else:
|
||||
self.use_triton_unified_attention = get_bool_env_var(
|
||||
@@ -2355,8 +2359,8 @@ class AiterAttnBackend(AttentionBackend):
|
||||
self.use_triton_unified_attention
|
||||
and self.use_sliding_window_kv_pool
|
||||
):
|
||||
token_to_kv_pool = forward_batch.token_to_kv_pool
|
||||
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
|
||||
token_to_kv_pool = self.token_to_kv_pool
|
||||
k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer(
|
||||
layer.layer_id
|
||||
)
|
||||
slot_mapping_swa = token_to_kv_pool.full_to_swa_index_mapping
|
||||
@@ -2380,9 +2384,9 @@ class AiterAttnBackend(AttentionBackend):
|
||||
v_scale=v_descale,
|
||||
)
|
||||
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:
|
||||
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
|
||||
)
|
||||
|
||||
@@ -2392,8 +2396,8 @@ class AiterAttnBackend(AttentionBackend):
|
||||
kv_indptr = self.forward_metadata.kv_indptr
|
||||
kv_indices = self.forward_metadata.kv_indices
|
||||
qo_indptr = self.forward_metadata.qo_indptr
|
||||
K_Buffer = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
V_Buffer = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
K_Buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||
V_Buffer = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||
kv_lora_rank = V_Buffer.shape[-1]
|
||||
qk_rope_head_dim = K_Buffer.shape[-1] - kv_lora_rank
|
||||
qk_nope_head_dim = k.shape[-1] - qk_rope_head_dim
|
||||
@@ -2646,7 +2650,7 @@ class AiterAttnBackend(AttentionBackend):
|
||||
self._use_unified_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
|
||||
)
|
||||
page_table = self.forward_metadata.kv_indices
|
||||
@@ -2705,8 +2709,8 @@ class AiterAttnBackend(AttentionBackend):
|
||||
k.contiguous(),
|
||||
v.contiguous(),
|
||||
o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
|
||||
forward_batch.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_key_buffer(layer.layer_id),
|
||||
self.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
self.forward_metadata.qo_indptr,
|
||||
self.forward_metadata.kv_indptr,
|
||||
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)
|
||||
|
||||
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)
|
||||
|
||||
bs0 = forward_batch.batch_size + 1
|
||||
|
||||
@@ -2798,10 +2800,8 @@ class AiterAttnBackend(AttentionBackend):
|
||||
# use standard set_kv_buffer, as they lack SWA-specific attributes
|
||||
# like full_to_swa_index_mapping.
|
||||
if self.use_triton_unified_attention and self.use_sliding_window_kv_pool:
|
||||
token_to_kv_pool = forward_batch.token_to_kv_pool
|
||||
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
|
||||
layer.layer_id
|
||||
)
|
||||
token_to_kv_pool = self.token_to_kv_pool
|
||||
k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
||||
slot_mapping_swa = token_to_kv_pool.full_to_swa_index_mapping
|
||||
|
||||
launch_reshape_and_cache_flash(
|
||||
@@ -2822,7 +2822,7 @@ class AiterAttnBackend(AttentionBackend):
|
||||
# [PATCH] FP8 non-SWA: use launch_reshape_and_cache_flash to
|
||||
# fuse bf16→fp8 cast + paged write in one Triton kernel,
|
||||
# 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)
|
||||
launch_reshape_and_cache_flash(
|
||||
k.view(-1, layer.tp_k_head_num, layer.qk_head_dim),
|
||||
@@ -2836,12 +2836,12 @@ class AiterAttnBackend(AttentionBackend):
|
||||
forward_batch.out_cache_loc,
|
||||
)
|
||||
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
|
||||
)
|
||||
|
||||
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_indptr = self.forward_metadata.work_indptr
|
||||
@@ -2878,9 +2878,7 @@ class AiterAttnBackend(AttentionBackend):
|
||||
else:
|
||||
self.logits_soft_cap = layer.logit_cap
|
||||
|
||||
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)
|
||||
|
||||
if layer.qk_head_dim != layer.v_head_dim:
|
||||
o = q.new_empty(
|
||||
@@ -3185,6 +3183,7 @@ class AiterMultiStepDraftBackend:
|
||||
)
|
||||
self.device = model_runner.device
|
||||
# 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.page_size = model_runner.server_args.page_size
|
||||
|
||||
@@ -3199,7 +3198,7 @@ class AiterMultiStepDraftBackend:
|
||||
(self.speculative_num_steps, num_seqs, self.topk)
|
||||
](
|
||||
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,
|
||||
kv_indices_buffer,
|
||||
self.kv_indptr,
|
||||
|
||||
@@ -241,14 +241,14 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
|
||||
assert v is not None
|
||||
if save_kv_cache:
|
||||
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,
|
||||
)
|
||||
else:
|
||||
forward_batch.token_to_kv_pool.set_kv_buffer(
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
k,
|
||||
@@ -269,7 +269,7 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
|
||||
q_nope = q_nope.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(
|
||||
q_nope=q_nope,
|
||||
|
||||
@@ -354,8 +354,15 @@ class DeepseekV4AttnBackend(
|
||||
self.page_size = model_runner.page_size
|
||||
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
|
||||
# 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]
|
||||
|
||||
assert isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool)
|
||||
@@ -378,6 +385,12 @@ class DeepseekV4AttnBackend(
|
||||
] = None
|
||||
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:
|
||||
pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True)
|
||||
return pin_tensor.to(self.device, non_blocking=True)
|
||||
@@ -667,7 +680,7 @@ class DeepseekV4AttnBackend(
|
||||
req_pool_indices = forward_batch.req_pool_indices
|
||||
seq_lens = forward_batch.seq_lens.to(torch.int32)
|
||||
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 seq_lens_cpu is not None
|
||||
@@ -960,7 +973,7 @@ class DeepseekV4AttnBackend(
|
||||
layer_id = layer.layer_id
|
||||
metadata = self.forward_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)
|
||||
|
||||
if isinstance(core_attn_metadata, DSV4AttnMetadata):
|
||||
|
||||
@@ -348,8 +348,15 @@ class DeepseekV4HipRadixBackend(
|
||||
self.page_size = model_runner.page_size
|
||||
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
|
||||
# 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]
|
||||
|
||||
assert isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool)
|
||||
@@ -372,6 +379,12 @@ class DeepseekV4HipRadixBackend(
|
||||
] = None
|
||||
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:
|
||||
pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True)
|
||||
return pin_tensor.to(self.device, non_blocking=True)
|
||||
@@ -661,7 +674,7 @@ class DeepseekV4HipRadixBackend(
|
||||
req_pool_indices = forward_batch.req_pool_indices
|
||||
seq_lens = forward_batch.seq_lens.to(torch.int32)
|
||||
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 seq_lens_cpu is not None
|
||||
@@ -954,7 +967,7 @@ class DeepseekV4HipRadixBackend(
|
||||
layer_id = layer.layer_id
|
||||
metadata = self.forward_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)
|
||||
|
||||
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.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_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
|
||||
|
||||
_use_ag_after_qlora = envs.SGLANG_USE_AG_AFTER_QLORA.get()
|
||||
@@ -449,9 +454,9 @@ class Indexer(MultiPlatformOp):
|
||||
metadata: BaseIndexerMetadata,
|
||||
) -> torch.Tensor:
|
||||
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
|
||||
if _is_hip:
|
||||
if _use_aiter_preshuffle:
|
||||
@@ -471,7 +476,7 @@ class Indexer(MultiPlatformOp):
|
||||
block_tables = metadata.get_page_table_64()
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
@@ -624,11 +629,11 @@ class Indexer(MultiPlatformOp):
|
||||
metadata: BaseIndexerMetadata,
|
||||
) -> torch.Tensor:
|
||||
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()
|
||||
|
||||
page_size = forward_batch.token_to_kv_pool.page_size
|
||||
page_size = get_token_to_kv_pool().page_size
|
||||
if _is_hip:
|
||||
if _use_aiter_preshuffle:
|
||||
assert (
|
||||
@@ -675,7 +680,7 @@ class Indexer(MultiPlatformOp):
|
||||
indexer_seq_lens_cpu = metadata.get_indexer_seq_len_cpu()
|
||||
seq_len_sum = torch.sum(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,
|
||||
metadata.get_indexer_seq_len(),
|
||||
block_tables,
|
||||
@@ -851,9 +856,9 @@ class Indexer(MultiPlatformOp):
|
||||
cp_index: List[Tuple[int, int, int]] = None,
|
||||
) -> torch.Tensor:
|
||||
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 len(weights.shape) == 3
|
||||
weights = weights.squeeze(-1)
|
||||
@@ -882,12 +887,12 @@ class Indexer(MultiPlatformOp):
|
||||
end_seq_position += pre_chunk_offset
|
||||
if offset == 0 and batch_idx != 0:
|
||||
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,
|
||||
end_seq_position,
|
||||
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,
|
||||
end_seq_position,
|
||||
block_tables[batch_idx],
|
||||
@@ -943,12 +948,12 @@ class Indexer(MultiPlatformOp):
|
||||
- forward_batch.extend_seq_lens_cpu[0]
|
||||
+ 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,
|
||||
kv_len,
|
||||
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,
|
||||
kv_len,
|
||||
block_tables[0],
|
||||
@@ -999,7 +1004,7 @@ class Indexer(MultiPlatformOp):
|
||||
if not _is_npu:
|
||||
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 len(weights.shape) == 3
|
||||
@@ -1011,7 +1016,7 @@ class Indexer(MultiPlatformOp):
|
||||
|
||||
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, :
|
||||
]
|
||||
strided_indices = torch.arange(
|
||||
@@ -1036,12 +1041,12 @@ class Indexer(MultiPlatformOp):
|
||||
weights_partial = weights[q_len_start:q_len_end]
|
||||
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,
|
||||
seq_len,
|
||||
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,
|
||||
seq_len,
|
||||
block_tables[i],
|
||||
@@ -1093,18 +1098,18 @@ class Indexer(MultiPlatformOp):
|
||||
and can_use_dsa_fused_store(
|
||||
key.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.
|
||||
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
|
||||
)
|
||||
fused_store_index_k_cache(
|
||||
key,
|
||||
buf,
|
||||
forward_batch.out_cache_loc,
|
||||
forward_batch.token_to_kv_pool.page_size,
|
||||
get_token_to_kv_pool().page_size,
|
||||
)
|
||||
return
|
||||
|
||||
@@ -1114,8 +1119,8 @@ class Indexer(MultiPlatformOp):
|
||||
# layout with page_size=1; the same kv_cache.view works for both cases
|
||||
# because page_size is 1 there.
|
||||
if _use_aiter:
|
||||
page_size = forward_batch.token_to_kv_pool.page_size
|
||||
buf = forward_batch.token_to_kv_pool.get_index_k_with_scale_buffer(
|
||||
page_size = get_token_to_kv_pool().page_size
|
||||
buf = get_token_to_kv_pool().get_index_k_with_scale_buffer(
|
||||
layer_id=layer_id
|
||||
)
|
||||
kv_cache = buf.view(-1, page_size, 132).view(fp8_dtype)
|
||||
@@ -1140,7 +1145,7 @@ class Indexer(MultiPlatformOp):
|
||||
if not out_loc.is_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,
|
||||
loc=out_loc,
|
||||
index_k=k_fp8,
|
||||
@@ -1175,15 +1180,13 @@ class Indexer(MultiPlatformOp):
|
||||
from sglang.srt.layers.attention.dsa.triton_kernel import act_quant
|
||||
|
||||
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
|
||||
# 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
|
||||
|
||||
metadata = forward_batch.attn_backend.get_indexer_metadata(
|
||||
layer_id, forward_batch
|
||||
)
|
||||
metadata = get_attn_backend().get_indexer_metadata(layer_id, forward_batch)
|
||||
|
||||
enable_dual_stream = (
|
||||
self.alt_stream is not None
|
||||
@@ -1405,12 +1408,10 @@ class Indexer(MultiPlatformOp):
|
||||
layer_scatter_modes=None,
|
||||
dynamic_scale: torch.Tensor = None,
|
||||
) -> torch.Tensor:
|
||||
if forward_batch.attn_backend.forward_metadata.seq_lens_cpu_int is None:
|
||||
actual_seq_lengths_kv = forward_batch.attn_backend.forward_metadata.seq_lens
|
||||
if get_attn_backend().forward_metadata.seq_lens_cpu_int is None:
|
||||
actual_seq_lengths_kv = get_attn_backend().forward_metadata.seq_lens
|
||||
else:
|
||||
actual_seq_lengths_kv = (
|
||||
forward_batch.attn_backend.forward_metadata.seq_lens_cpu_int
|
||||
)
|
||||
actual_seq_lengths_kv = get_attn_backend().forward_metadata.seq_lens_cpu_int
|
||||
is_prefill = (
|
||||
forward_batch.forward_mode.is_extend()
|
||||
and not forward_batch.forward_mode.is_draft_extend_v2()
|
||||
@@ -1558,7 +1559,7 @@ class Indexer(MultiPlatformOp):
|
||||
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
|
||||
)
|
||||
if is_prefill:
|
||||
@@ -1566,7 +1567,7 @@ class Indexer(MultiPlatformOp):
|
||||
self.dsa_enable_prefill_cp
|
||||
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_next_tensor,
|
||||
)
|
||||
@@ -1579,34 +1580,32 @@ class Indexer(MultiPlatformOp):
|
||||
forward_batch.attn_cp_metadata.kv_len_next_tensor
|
||||
+ 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_next_tensor,
|
||||
)
|
||||
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_next_tensor,
|
||||
)
|
||||
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 = (
|
||||
forward_batch.attn_backend.forward_metadata.actual_seq_lengths_kv
|
||||
get_attn_backend().forward_metadata.actual_seq_lengths_kv
|
||||
)
|
||||
else:
|
||||
actual_seq_lengths_kv = forward_batch.seq_lens
|
||||
actual_seq_lengths_q = forward_batch.extend_seq_lens.cumsum(dim=0)
|
||||
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 (
|
||||
forward_batch.forward_mode.is_draft_extend_v2()
|
||||
or forward_batch.forward_mode.is_target_verify()
|
||||
or forward_batch.forward_mode.is_draft_extend()
|
||||
):
|
||||
num_draft_tokens = (
|
||||
forward_batch.attn_backend.speculative_num_draft_tokens
|
||||
)
|
||||
num_draft_tokens = get_attn_backend().speculative_num_draft_tokens
|
||||
actual_seq_lengths_q = torch.arange(
|
||||
num_draft_tokens,
|
||||
num_draft_tokens + bs,
|
||||
@@ -1622,10 +1621,10 @@ class Indexer(MultiPlatformOp):
|
||||
)
|
||||
else:
|
||||
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:
|
||||
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
|
||||
):
|
||||
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 (
|
||||
is_prefill
|
||||
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
|
||||
|
||||
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.use_mha: bool = False
|
||||
@@ -392,6 +400,12 @@ class DeepseekSparseAttnBackend(
|
||||
else:
|
||||
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:
|
||||
if l > len(self._arange_buf):
|
||||
next_pow_of_2 = 1 << (l - 1).bit_length()
|
||||
@@ -425,7 +439,7 @@ class DeepseekSparseAttnBackend(
|
||||
assert forward_batch.seq_lens_cpu is not None
|
||||
max_seqlen_k = int(forward_batch.seq_lens_cpu.max().item() + draft_token_num)
|
||||
# [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
|
||||
]
|
||||
|
||||
@@ -580,8 +594,7 @@ class DeepseekSparseAttnBackend(
|
||||
|
||||
# Check if MHA FP8 dequantization is needed
|
||||
mha_dequantize_needed = (
|
||||
self.use_mha
|
||||
and forward_batch.token_to_kv_pool.dtype == torch.float8_e4m3fn
|
||||
self.use_mha and self.token_to_kv_pool.dtype == torch.float8_e4m3fn
|
||||
)
|
||||
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
|
||||
# This is likely to be triggered by PP with high kv reuse & parallelism
|
||||
kv_cache_capacity = (
|
||||
forward_batch.token_to_kv_pool.size
|
||||
+ forward_batch.token_to_kv_pool.page_size
|
||||
self.token_to_kv_pool.size + self.token_to_kv_pool.page_size
|
||||
)
|
||||
if forward_batch.seq_lens_sum > kv_cache_capacity:
|
||||
max_idx = page_table_1_flattened.max().item()
|
||||
@@ -1380,7 +1392,7 @@ class DeepseekSparseAttnBackend(
|
||||
if not layer.is_cross_attention
|
||||
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,
|
||||
cache_loc,
|
||||
k,
|
||||
@@ -1405,7 +1417,7 @@ class DeepseekSparseAttnBackend(
|
||||
|
||||
# Do absorbed multi-latent attention (MLA path)
|
||||
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:
|
||||
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
|
||||
if forward_batch.hisparse_coordinator is not None:
|
||||
page_table_1 = (
|
||||
forward_batch.token_to_kv_pool.translate_loc_to_hisparse_device(
|
||||
page_table_1
|
||||
)
|
||||
if self.hisparse_coordinator is not None:
|
||||
page_table_1 = self.token_to_kv_pool.translate_loc_to_hisparse_device(
|
||||
page_table_1
|
||||
)
|
||||
|
||||
if dsa_impl == "tilelang":
|
||||
@@ -1580,7 +1590,7 @@ class DeepseekSparseAttnBackend(
|
||||
if not layer.is_cross_attention
|
||||
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,
|
||||
cache_loc,
|
||||
k,
|
||||
@@ -1588,7 +1598,7 @@ class DeepseekSparseAttnBackend(
|
||||
)
|
||||
|
||||
# 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:
|
||||
q_nope = q.view(-1, layer.tp_q_head_num, layer.v_head_dim)
|
||||
q_rope = q_rope.view(
|
||||
@@ -1609,8 +1619,8 @@ class DeepseekSparseAttnBackend(
|
||||
if topk_indices is not None:
|
||||
topk_indices = self._pad_topk_indices(topk_indices, q_nope.shape[0])
|
||||
|
||||
if forward_batch.hisparse_coordinator is not None:
|
||||
page_table_1 = forward_batch.hisparse_coordinator.swap_in_selected_pages(
|
||||
if self.hisparse_coordinator is not None:
|
||||
page_table_1 = self.hisparse_coordinator.swap_in_selected_pages(
|
||||
forward_batch.req_pool_indices,
|
||||
forward_batch.seq_lens,
|
||||
topk_indices,
|
||||
@@ -2105,11 +2115,9 @@ class DeepseekSparseAttnBackend(
|
||||
if not layer.is_cross_attention
|
||||
else forward_batch.encoder_out_cache_loc
|
||||
)
|
||||
forward_batch.token_to_kv_pool.set_mla_kv_buffer(
|
||||
layer, cache_loc, k, k_rope
|
||||
)
|
||||
self.token_to_kv_pool.set_mla_kv_buffer(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)
|
||||
|
||||
if merge_query:
|
||||
@@ -2221,12 +2229,11 @@ class DeepseekSparseAttnBackend(
|
||||
) # SM90/SM100 only
|
||||
and max_kv_len
|
||||
<= envs.SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD.get() # Short enough for MHA
|
||||
and forward_batch.token_to_kv_pool.dtype
|
||||
in [torch.bfloat16, torch.float8_e4m3fn]
|
||||
and self.token_to_kv_pool.dtype in [torch.bfloat16, torch.float8_e4m3fn]
|
||||
and sum_seq_lens
|
||||
<= forward_batch.get_max_chunk_capacity() # Fits in chunk
|
||||
and (not is_dsa_enable_prefill_cp()) # CP not enabled
|
||||
and (forward_batch.hisparse_coordinator is None)
|
||||
and (self.hisparse_coordinator is None)
|
||||
)
|
||||
else:
|
||||
self.use_mha = False # Decode/verify always use MLA
|
||||
@@ -2272,7 +2279,7 @@ class DeepseekSparseAttnBackend(
|
||||
self, layer_id: int, forward_batch: ForwardBatch
|
||||
) -> DSAIndexerMetadata:
|
||||
force_unfused = (
|
||||
forward_batch.hisparse_coordinator is not None
|
||||
self.hisparse_coordinator is not None
|
||||
and forward_batch.forward_mode.is_decode_or_idle()
|
||||
)
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import (
|
||||
DeepseekV4HipRadixBackend,
|
||||
)
|
||||
@@ -99,24 +100,26 @@ class CompressorHip(_CompressorBase):
|
||||
def use_hip_fused_compress(self) -> bool:
|
||||
return envs.SGLANG_OPT_USE_FUSED_COMPRESS.get()
|
||||
|
||||
def _get_states(self, forward_batch: ForwardBatch) -> KVAndScore:
|
||||
token_to_kv_pool = forward_batch.token_to_kv_pool
|
||||
def _get_states(
|
||||
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)
|
||||
if self.is_in_indexer:
|
||||
return token_to_kv_pool.get_indexer_compress_states(self.layer_id)
|
||||
else:
|
||||
return token_to_kv_pool.get_attention_compress_states(self.layer_id)
|
||||
|
||||
def _get_state_pool(self, forward_batch: ForwardBatch) -> CompressStatePool:
|
||||
token_to_kv_pool = forward_batch.token_to_kv_pool
|
||||
def _get_state_pool(self, attn_backend: AttentionBackend) -> CompressStatePool:
|
||||
token_to_kv_pool = attn_backend.token_to_kv_pool
|
||||
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
|
||||
if self.is_in_indexer:
|
||||
ret = token_to_kv_pool.get_indexer_compress_states(self.layer_id)
|
||||
else:
|
||||
ret = token_to_kv_pool.get_attention_compress_states(self.layer_id)
|
||||
|
||||
assert isinstance(ret, CompressStatePool)
|
||||
|
||||
return ret
|
||||
|
||||
def overlap_transform(self, tensor: torch.Tensor, fill_value: Any) -> torch.Tensor:
|
||||
@@ -155,18 +158,19 @@ class CompressorHip(_CompressorBase):
|
||||
self,
|
||||
kv_and_scores: KVAndScore,
|
||||
forward_batch: ForwardBatch,
|
||||
attn_backend: AttentionBackend,
|
||||
):
|
||||
backend = forward_batch.attn_backend
|
||||
backend = attn_backend
|
||||
if TYPE_CHECKING:
|
||||
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)
|
||||
|
||||
state_pool = self._get_state_pool(forward_batch)
|
||||
state_pool = self._get_state_pool(backend)
|
||||
prefix_lens = forward_batch.extend_prefix_lens_cpu
|
||||
extend_lens = forward_batch.extend_seq_lens_cpu
|
||||
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 extend_lens is not None and prefix_lens is not None
|
||||
@@ -289,18 +293,19 @@ class CompressorHip(_CompressorBase):
|
||||
self,
|
||||
kv_and_scores: KVAndScore,
|
||||
forward_batch: ForwardBatch,
|
||||
attn_backend: AttentionBackend,
|
||||
):
|
||||
"""Paged and cudagraph compatible version of compress_decode"""
|
||||
assert self.ape_converted
|
||||
state_pool = self._get_state_pool(forward_batch)
|
||||
token_to_kv_pool = forward_batch.token_to_kv_pool
|
||||
state_pool = self._get_state_pool(attn_backend)
|
||||
token_to_kv_pool = attn_backend.token_to_kv_pool
|
||||
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
|
||||
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
|
||||
|
||||
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)
|
||||
seq_lens_2d = seq_lens[:, None] + offsets[None, :]
|
||||
seq_lens = seq_lens_2d.view(-1)
|
||||
@@ -378,11 +383,12 @@ class CompressorHip(_CompressorBase):
|
||||
self,
|
||||
kv_score: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
attn_backend: AttentionBackend,
|
||||
) -> torch.Tensor:
|
||||
backend = forward_batch.attn_backend
|
||||
backend = attn_backend
|
||||
if TYPE_CHECKING:
|
||||
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
|
||||
|
||||
return backend.forward_compress(
|
||||
@@ -402,9 +408,12 @@ class CompressorHip(_CompressorBase):
|
||||
self,
|
||||
kv_score: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
attn_backend: AttentionBackend,
|
||||
) -> torch.Tensor:
|
||||
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_extend = self.compress_extend_paged
|
||||
@@ -420,11 +429,13 @@ class CompressorHip(_CompressorBase):
|
||||
result = self.compress_decode(
|
||||
kv_and_scores=kv_and_scores,
|
||||
forward_batch=forward_batch,
|
||||
attn_backend=attn_backend,
|
||||
)
|
||||
elif forward_batch.forward_mode.is_extend():
|
||||
result = self.compress_extend(
|
||||
kv_and_scores=kv_and_scores,
|
||||
forward_batch=forward_batch,
|
||||
attn_backend=attn_backend,
|
||||
)
|
||||
else:
|
||||
msg = f"Forward mode {forward_batch.forward_mode} not supported in Compressor."
|
||||
@@ -445,11 +456,17 @@ class CompressorHip(_CompressorBase):
|
||||
setattr(forward_batch, attr, 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():
|
||||
assert x.shape[0] == 0
|
||||
return x.new_empty(0, self.head_dim)
|
||||
|
||||
kv_score = self.compute_kv_score(x, forward_batch)
|
||||
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
|
||||
|
||||
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.rotary_embedding import RotaryEmbedding
|
||||
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
|
||||
# (e.g. 1.6T layer 0 has compress_ratio=128 and needs cX_compress_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:
|
||||
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
|
||||
out_loc = (
|
||||
core_metadata.c4_out_loc
|
||||
@@ -154,11 +155,11 @@ class CompressorBackendMixin:
|
||||
assert is_overlap_compress(compressor.ratio)
|
||||
# PREP_IN_CG lazy upgrade (see forward_core_compressor for rationale).
|
||||
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:
|
||||
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():
|
||||
token_to_kv_pool.set_index_k_fused(
|
||||
layer_id=layer_id,
|
||||
@@ -339,20 +340,16 @@ class Compressor(nn.Module):
|
||||
ape = torch.cat([ape[0], ape[1]], dim=0)
|
||||
self.ape.data.copy_(ape.view(self.ratio, -1))
|
||||
|
||||
# NOTE: used by v2 compressor backend
|
||||
def get_state_pool(self, forward_batch: ForwardBatch) -> CompressStatePool:
|
||||
token_to_kv_pool = forward_batch.token_to_kv_pool
|
||||
def get_state_pool(self, attn_backend: AttentionBackend) -> CompressStatePool:
|
||||
token_to_kv_pool = attn_backend.token_to_kv_pool
|
||||
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
|
||||
if self.is_in_indexer:
|
||||
ret = token_to_kv_pool.get_indexer_compress_states(self.layer_id)
|
||||
else:
|
||||
ret = token_to_kv_pool.get_attention_compress_states(self.layer_id)
|
||||
|
||||
assert isinstance(ret, CompressStatePool)
|
||||
|
||||
return ret
|
||||
|
||||
# NOTE: used by v2 compressor backend
|
||||
def compute_kv_score(self, x: torch.Tensor, forward_batch: ForwardBatch):
|
||||
kv_score = linear_bf16_fp32(x, self.wkv_gate.weight)
|
||||
|
||||
@@ -366,19 +363,22 @@ class Compressor(nn.Module):
|
||||
)
|
||||
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():
|
||||
assert x.shape[0] == 0
|
||||
return x.new_empty(0, self.head_dim)
|
||||
|
||||
kv_score = self.compute_kv_score(x, forward_batch)
|
||||
|
||||
backend = forward_batch.attn_backend
|
||||
if TYPE_CHECKING:
|
||||
assert isinstance(backend, DeepseekV4AttnBackend)
|
||||
kv_score_buffer = self.get_state_pool(forward_batch)
|
||||
kv_score_buffer = kv_score_buffer.kv_score_buffer.kv_score
|
||||
return backend.forward_compress(
|
||||
assert isinstance(attn_backend, DeepseekV4AttnBackend)
|
||||
kv_score_buffer = self.get_state_pool(attn_backend).kv_score_buffer.kv_score
|
||||
return attn_backend.forward_compress(
|
||||
kv_score_buffer=kv_score_buffer,
|
||||
kv_score_input=kv_score,
|
||||
ape=self.ape.view(-1, self.head_dim),
|
||||
|
||||
@@ -106,10 +106,10 @@ class CompressorBackendMixin:
|
||||
return
|
||||
|
||||
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)
|
||||
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)
|
||||
if compressor.is_in_indexer:
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.attention.dsv4.compressor import (
|
||||
CompressorBackendMixin,
|
||||
)
|
||||
@@ -321,7 +322,7 @@ class C4IndexerBackendMixin:
|
||||
# PREP_IN_CG lazy upgrade: this runs from MQALayer._forward_prepare,
|
||||
# before attn_backend.forward() would trigger the upgrade.
|
||||
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:
|
||||
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
|
||||
@@ -398,7 +399,7 @@ class C4IndexerBackendMixin:
|
||||
indexer_capturer = get_global_indexer_capturer()
|
||||
capture_enabled = indexer_capturer is not None
|
||||
|
||||
hisparse_coordinator = forward_batch.hisparse_coordinator
|
||||
hisparse_coordinator = self.hisparse_coordinator
|
||||
hisparse_decode = (
|
||||
hisparse_coordinator is not None and forward_batch.forward_mode.is_decode()
|
||||
)
|
||||
@@ -541,10 +542,11 @@ class C4Indexer(nn.Module):
|
||||
x: torch.Tensor,
|
||||
q_lora: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
attn_backend: AttentionBackend,
|
||||
enable_multi_stream: bool = False,
|
||||
q_lora_ready: Optional[torch.cuda.Event] = None,
|
||||
) -> None:
|
||||
return forward_batch.attn_backend.forward_c4_indexer(
|
||||
return attn_backend.forward_c4_indexer(
|
||||
x=x,
|
||||
q_lora=q_lora,
|
||||
forward_batch=forward_batch,
|
||||
|
||||
@@ -117,6 +117,10 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
|
||||
)
|
||||
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.kv_cache_dtype = model_runner.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 = 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
|
||||
]
|
||||
# Convert the block table to a strided format.
|
||||
@@ -346,9 +350,7 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
|
||||
assert current_end <= self.max_context_len
|
||||
|
||||
# Do multi-head attention
|
||||
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
|
||||
layer.layer_id
|
||||
)
|
||||
key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
||||
key_cache = key_cache.view(
|
||||
-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 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,
|
||||
key,
|
||||
@@ -442,9 +444,7 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
|
||||
key = k.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(
|
||||
layer.layer_id
|
||||
)
|
||||
key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
||||
key_cache = key_cache.view(
|
||||
-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 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,
|
||||
key,
|
||||
|
||||
@@ -126,6 +126,10 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
self.device = model_runner.device
|
||||
self.decode_cuda_graph_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.kv_cache_dtype = model_runner.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)
|
||||
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.speculative_num_steps = speculative_num_steps
|
||||
@@ -295,7 +297,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
),
|
||||
(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
|
||||
]
|
||||
else:
|
||||
@@ -315,7 +317,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
),
|
||||
(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
|
||||
]
|
||||
metadata_expand = FlashAttentionMetadata()
|
||||
@@ -358,7 +360,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
metadata.cu_seqlens_k = torch.nn.functional.pad(
|
||||
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
|
||||
]
|
||||
# Precompute FA3 scheduler metadata to avoid per-layer
|
||||
@@ -394,7 +396,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
),
|
||||
(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
|
||||
]
|
||||
|
||||
@@ -416,7 +418,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
),
|
||||
(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
|
||||
]
|
||||
|
||||
@@ -477,7 +479,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
)
|
||||
_, sort_order = torch.sort(keys, dim=1)
|
||||
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, :
|
||||
]
|
||||
.gather(1, cols)
|
||||
@@ -506,7 +508,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
metadata.cu_seqlens_k = torch.nn.functional.pad(
|
||||
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
|
||||
]
|
||||
|
||||
@@ -538,12 +540,12 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
(1, 0),
|
||||
)
|
||||
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
|
||||
]
|
||||
|
||||
# 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,
|
||||
metadata.encoder_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
|
||||
)
|
||||
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
|
||||
)
|
||||
else:
|
||||
forward_batch.token_to_kv_pool.set_mla_kv_buffer(
|
||||
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
k,
|
||||
@@ -746,9 +748,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
# Use Flash Attention for prefill
|
||||
if not self.use_mla:
|
||||
# Do multi-head attention
|
||||
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
|
||||
layer.layer_id
|
||||
)
|
||||
key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
||||
|
||||
key_cache = key_cache.view(
|
||||
-1, self.page_size, layer.tp_k_head_num, layer.head_dim
|
||||
@@ -950,9 +950,9 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
else:
|
||||
assert self.fa_impl_ver == 3, "Only FA3 support here"
|
||||
# Do absorbed multi-latent attention
|
||||
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer(
|
||||
layer.layer_id
|
||||
).to(q.dtype)
|
||||
kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
|
||||
q.dtype
|
||||
)
|
||||
k_rope = kv_cache[:, :, layer.v_head_dim :]
|
||||
c_kv = kv_cache[:, :, : layer.v_head_dim]
|
||||
k_rope_cache = k_rope.view(
|
||||
@@ -1050,11 +1050,11 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
else forward_batch.encoder_out_cache_loc
|
||||
)
|
||||
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
|
||||
)
|
||||
else:
|
||||
forward_batch.token_to_kv_pool.set_mla_kv_buffer(
|
||||
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
k,
|
||||
@@ -1113,9 +1113,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
if not self.use_mla:
|
||||
# Do multi-head attention
|
||||
|
||||
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
|
||||
layer.layer_id
|
||||
)
|
||||
key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
||||
key_cache = key_cache.view(
|
||||
-1, self.page_size, layer.tp_k_head_num, layer.head_dim
|
||||
)
|
||||
@@ -1248,9 +1246,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
o = result
|
||||
else:
|
||||
# Do absorbed multi-latent attention
|
||||
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
|
||||
q.dtype
|
||||
)
|
||||
kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(q.dtype)
|
||||
k_rope = kv_cache[:, :, layer.v_head_dim :]
|
||||
c_kv = kv_cache[:, :, : layer.v_head_dim]
|
||||
k_rope_cache = k_rope.view(
|
||||
|
||||
@@ -126,7 +126,8 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
self.prefill_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
|
||||
|
||||
# FIXME: remove dllm workarounds from flashinfer
|
||||
@@ -802,7 +803,7 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
if k is not None:
|
||||
assert v is not None
|
||||
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
|
||||
)
|
||||
|
||||
@@ -812,7 +813,7 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
)
|
||||
o = prefill_wrapper_paged.forward(
|
||||
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,
|
||||
sm_scale=layer.scaling,
|
||||
# Disable sliding window attention for multi-item scoring:
|
||||
@@ -836,12 +837,12 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
)
|
||||
else:
|
||||
# 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
|
||||
# IQuestLoopCoder path uses token_to_kv_pool as the KV source).
|
||||
if k is None and v is None:
|
||||
k = forward_batch.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]
|
||||
k = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)[0]
|
||||
v = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)[1]
|
||||
causal = True
|
||||
if (
|
||||
layer.is_cross_attention
|
||||
@@ -875,7 +876,7 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
)
|
||||
o2, s2 = prefill_wrapper_paged.forward_return_lse(
|
||||
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,
|
||||
sm_scale=layer.scaling,
|
||||
logits_soft_cap=logits_soft_cap,
|
||||
@@ -884,7 +885,7 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
o, _ = merge_state(o1, s1, o2, s2)
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
@@ -912,14 +913,14 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
if k is not None:
|
||||
assert v is not None
|
||||
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
|
||||
)
|
||||
|
||||
# Call the wrapped function
|
||||
o = decode_wrapper.forward(
|
||||
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,
|
||||
logits_soft_cap=layer.logit_cap,
|
||||
# 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
|
||||
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(
|
||||
self,
|
||||
@@ -1562,7 +1564,7 @@ class FlashInferMultiStepDraftBackend:
|
||||
(self.speculative_num_steps, num_seqs, self.topk)
|
||||
](
|
||||
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,
|
||||
kv_indices_buffer,
|
||||
self.kv_indptr,
|
||||
|
||||
@@ -204,6 +204,10 @@ class FlashInferMLAAttnBackend(AttentionBackend):
|
||||
self.max_context_len = model_runner.model_config.context_len
|
||||
self.device = model_runner.device
|
||||
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 = (
|
||||
not skip_prefill
|
||||
and get_global_server_args().disaggregation_mode != "decode"
|
||||
@@ -544,11 +548,9 @@ class FlashInferMLAAttnBackend(AttentionBackend):
|
||||
assert v is not None
|
||||
if save_kv_cache:
|
||||
if k_rope is not None:
|
||||
forward_batch.token_to_kv_pool.set_mla_kv_buffer(
|
||||
layer, cache_loc, k, k_rope
|
||||
)
|
||||
self.token_to_kv_pool.set_mla_kv_buffer(layer, cache_loc, k, k_rope)
|
||||
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:
|
||||
q = q.view(-1, layer.tp_q_head_num, layer.v_head_dim)
|
||||
q_rope = q_rope.view(
|
||||
@@ -572,9 +574,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
|
||||
)
|
||||
else:
|
||||
# mla paged prefill
|
||||
k_buf = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
|
||||
q.dtype
|
||||
)
|
||||
k_buf = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(q.dtype)
|
||||
if q_rope is None:
|
||||
qall = q.view(-1, layer.tp_q_head_num, layer.head_dim)
|
||||
q, q_rope = (
|
||||
@@ -611,14 +611,14 @@ class FlashInferMLAAttnBackend(AttentionBackend):
|
||||
assert v is not None
|
||||
if save_kv_cache:
|
||||
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,
|
||||
)
|
||||
else:
|
||||
forward_batch.token_to_kv_pool.set_kv_buffer(
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
k,
|
||||
@@ -636,9 +636,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
|
||||
q_nope = 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(
|
||||
q.dtype
|
||||
)
|
||||
k_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(q.dtype)
|
||||
|
||||
o = q_nope.new_empty(q_nope.shape)
|
||||
# Direct call to run without the wrapper
|
||||
@@ -944,6 +942,7 @@ class FlashInferMLAMultiStepDraftBackend:
|
||||
self.max_context_len = self.attn_backends[0].max_context_len
|
||||
|
||||
# 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.page_size = model_runner.server_args.page_size
|
||||
|
||||
@@ -961,7 +960,7 @@ class FlashInferMLAMultiStepDraftBackend:
|
||||
(self.speculative_num_steps, num_seqs, self.topk)
|
||||
](
|
||||
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,
|
||||
kv_indices_buffer,
|
||||
self.kv_indptr,
|
||||
|
||||
@@ -410,14 +410,14 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
||||
if k is not None:
|
||||
assert v is not None
|
||||
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,
|
||||
)
|
||||
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)
|
||||
if self.is_fp8_kvcache:
|
||||
@@ -489,10 +489,10 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
||||
if k is not None:
|
||||
assert v is not None
|
||||
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
|
||||
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)
|
||||
if self.is_fp8_kvcache:
|
||||
|
||||
@@ -23,6 +23,8 @@ class HybridAttnBackend(AttentionBackend):
|
||||
self.prefill_backend = prefill_backend
|
||||
self.decode_backend = decode_backend
|
||||
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:
|
||||
"""
|
||||
|
||||
@@ -143,6 +143,7 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
||||
self.is_draft_worker = model_runner.is_draft_worker
|
||||
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.state_indices_list = []
|
||||
self.query_start_loc_list = []
|
||||
@@ -763,6 +764,9 @@ class HybridLinearAttnBackend(AttentionBackend):
|
||||
self.full_attn_backend = full_attn_backend
|
||||
self.linear_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(
|
||||
self, layer: Optional[RadixAttention], layer_id: Optional[int] = None
|
||||
|
||||
@@ -19,6 +19,10 @@ class IntelAMXAttnBackend(AttentionBackend):
|
||||
super().__init__()
|
||||
self.forward_metadata = None
|
||||
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 = (
|
||||
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
|
||||
)
|
||||
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
|
||||
self.extend_attention_fwd(
|
||||
@@ -113,9 +117,9 @@ class IntelAMXAttnBackend(AttentionBackend):
|
||||
k,
|
||||
v,
|
||||
o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
|
||||
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id),
|
||||
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
forward_batch.req_to_token_pool.req_to_token,
|
||||
self.token_to_kv_pool.get_key_buffer(layer.layer_id),
|
||||
self.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
self.req_to_token_pool.req_to_token,
|
||||
forward_batch.req_pool_indices,
|
||||
forward_batch.seq_lens,
|
||||
forward_batch.extend_seq_lens,
|
||||
@@ -152,14 +156,14 @@ class IntelAMXAttnBackend(AttentionBackend):
|
||||
)
|
||||
self.decode_attention_fwd(
|
||||
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
|
||||
forward_batch.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_key_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),
|
||||
k,
|
||||
v,
|
||||
cache_loc,
|
||||
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.seq_lens,
|
||||
layer.scaling,
|
||||
|
||||
@@ -15,6 +15,10 @@ class TboAttnBackend(AttentionBackend):
|
||||
super().__init__()
|
||||
self.primary = primary
|
||||
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
|
||||
def init_new(cls, creator: Callable[[], AttentionBackend]):
|
||||
|
||||
@@ -249,7 +249,7 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
||||
# reproduces the original [tokens, 1, qk_rope] latent layout.
|
||||
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 :]
|
||||
forward_batch.token_to_kv_pool.set_mla_kv_buffer(
|
||||
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||
layer.attn_mha,
|
||||
forward_batch.out_cache_loc,
|
||||
kv_a_fp8.unsqueeze(1),
|
||||
|
||||
@@ -19,6 +19,10 @@ class TorchFlexAttnBackend(AttentionBackend):
|
||||
super().__init__()
|
||||
self.forward_metadata = None
|
||||
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)
|
||||
torch._dynamo.config.cache_size_limit = 1024
|
||||
torch._dynamo.config.accumulated_cache_size_limit = 1024
|
||||
@@ -248,7 +252,7 @@ class TorchFlexAttnBackend(AttentionBackend):
|
||||
o = torch.empty_like(q)
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
@@ -266,9 +270,9 @@ class TorchFlexAttnBackend(AttentionBackend):
|
||||
self._run_flex_forward_extend(
|
||||
q_,
|
||||
o_,
|
||||
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id),
|
||||
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
forward_batch.req_to_token_pool.req_to_token,
|
||||
self.token_to_kv_pool.get_key_buffer(layer.layer_id),
|
||||
self.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
self.req_to_token_pool.req_to_token,
|
||||
forward_batch.req_pool_indices,
|
||||
forward_batch.seq_lens,
|
||||
forward_batch.extend_prefix_lens,
|
||||
@@ -298,7 +302,7 @@ class TorchFlexAttnBackend(AttentionBackend):
|
||||
o = torch.empty_like(q)
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
@@ -309,9 +313,9 @@ class TorchFlexAttnBackend(AttentionBackend):
|
||||
self._run_flex_forward_decode(
|
||||
q_,
|
||||
o_,
|
||||
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id),
|
||||
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
forward_batch.req_to_token_pool.req_to_token,
|
||||
self.token_to_kv_pool.get_key_buffer(layer.layer_id),
|
||||
self.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
self.req_to_token_pool.req_to_token,
|
||||
forward_batch.req_pool_indices,
|
||||
forward_batch.seq_lens,
|
||||
scaling=layer.scaling,
|
||||
|
||||
@@ -19,6 +19,10 @@ class TorchNativeAttnBackend(AttentionBackend):
|
||||
super().__init__()
|
||||
self.forward_metadata = None
|
||||
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):
|
||||
"""Init the metadata for a forward pass."""
|
||||
@@ -235,7 +239,7 @@ class TorchNativeAttnBackend(AttentionBackend):
|
||||
cache_loc = forward_batch.out_cache_loc
|
||||
|
||||
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
|
||||
|
||||
@@ -249,9 +253,9 @@ class TorchNativeAttnBackend(AttentionBackend):
|
||||
self._run_sdpa_forward_extend(
|
||||
q_,
|
||||
o_,
|
||||
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id),
|
||||
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
forward_batch.req_to_token_pool.req_to_token,
|
||||
self.token_to_kv_pool.get_key_buffer(layer.layer_id),
|
||||
self.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
self.req_to_token_pool.req_to_token,
|
||||
forward_batch.req_pool_indices,
|
||||
forward_batch.seq_lens,
|
||||
forward_batch.extend_prefix_lens,
|
||||
@@ -292,9 +296,8 @@ class TorchNativeAttnBackend(AttentionBackend):
|
||||
else:
|
||||
cache_loc = forward_batch.out_cache_loc
|
||||
|
||||
if save_kv_cache:
|
||||
if k is not None and v is not None:
|
||||
forward_batch.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
|
||||
if save_kv_cache and k is not None and v is not None:
|
||||
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
|
||||
|
||||
@@ -304,9 +307,9 @@ class TorchNativeAttnBackend(AttentionBackend):
|
||||
self._run_sdpa_forward_decode(
|
||||
q_,
|
||||
o_,
|
||||
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id),
|
||||
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
forward_batch.req_to_token_pool.req_to_token,
|
||||
self.token_to_kv_pool.get_key_buffer(layer.layer_id),
|
||||
self.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
self.req_to_token_pool.req_to_token,
|
||||
forward_batch.req_pool_indices,
|
||||
forward_batch.seq_lens,
|
||||
forward_batch.encoder_lens,
|
||||
|
||||
@@ -105,6 +105,10 @@ class TritonAttnBackend(AttentionBackend):
|
||||
self.skip_prefill = skip_prefill
|
||||
max_bs = model_runner.req_to_token_pool.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.token_to_kv_pool_allocator = model_runner.token_to_kv_pool_allocator
|
||||
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
|
||||
@@ -904,7 +908,7 @@ class TritonAttnBackend(AttentionBackend):
|
||||
o = torch.empty_like(q)
|
||||
|
||||
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
|
||||
if isinstance(pool, SWAKVPool) and pool.layers_mapping[layer.layer_id][1]:
|
||||
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)
|
||||
if save_kv_cache:
|
||||
if layer.k_scale is None:
|
||||
forward_batch.token_to_kv_pool.set_kv_buffer(
|
||||
self.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
forward_batch.out_cache_loc,
|
||||
k,
|
||||
@@ -928,14 +932,14 @@ class TritonAttnBackend(AttentionBackend):
|
||||
# doesn't accept scale parameters. Clone to protect k from mutation
|
||||
# since it's used later in the attention kernel.
|
||||
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,
|
||||
forward_batch.out_cache_loc,
|
||||
k_scaled,
|
||||
v,
|
||||
)
|
||||
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.clone(), # cloned to protect k,v from in-place mutation in set_kv_buffer
|
||||
@@ -989,8 +993,8 @@ class TritonAttnBackend(AttentionBackend):
|
||||
k.contiguous(),
|
||||
v.contiguous(),
|
||||
o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
|
||||
forward_batch.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_key_buffer(layer.layer_id),
|
||||
self.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
self.forward_metadata.qo_indptr,
|
||||
kv_indptr,
|
||||
kv_indices,
|
||||
@@ -1058,7 +1062,7 @@ class TritonAttnBackend(AttentionBackend):
|
||||
window_start_pos = None
|
||||
|
||||
extend_kv_indices = forward_batch.out_cache_loc
|
||||
pool = forward_batch.token_to_kv_pool
|
||||
pool = self.token_to_kv_pool
|
||||
if (
|
||||
layer.sliding_window_size is not None
|
||||
and layer.sliding_window_size > -1
|
||||
@@ -1124,8 +1128,8 @@ class TritonAttnBackend(AttentionBackend):
|
||||
self.extend_attention_fwd_unified(
|
||||
q.view(-1, layer.tp_q_head_num, layer.qk_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),
|
||||
forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
self.token_to_kv_pool.get_key_buffer(layer.layer_id),
|
||||
self.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
k_descale,
|
||||
v_descale,
|
||||
self.forward_metadata.qo_indptr,
|
||||
@@ -1174,14 +1178,14 @@ class TritonAttnBackend(AttentionBackend):
|
||||
# MLATokenToKVPool doesn't accept scale parameters; k is unused
|
||||
# after this point in decode, so scale in place.
|
||||
k.div_(layer.k_scale)
|
||||
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,
|
||||
)
|
||||
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,
|
||||
@@ -1216,8 +1220,8 @@ class TritonAttnBackend(AttentionBackend):
|
||||
|
||||
self.decode_attention_fwd(
|
||||
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
|
||||
forward_batch.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_key_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),
|
||||
kv_indptr,
|
||||
kv_indices,
|
||||
@@ -1275,6 +1279,7 @@ class TritonMultiStepDraftBackend:
|
||||
)
|
||||
self.device = model_runner.device
|
||||
# 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.page_size = model_runner.server_args.page_size
|
||||
|
||||
@@ -1295,7 +1300,7 @@ class TritonMultiStepDraftBackend:
|
||||
(self.speculative_num_steps, num_seqs, self.topk)
|
||||
](
|
||||
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,
|
||||
kv_indices_buffer,
|
||||
self.kv_indptr,
|
||||
|
||||
@@ -558,7 +558,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
cache_loc = self._get_layer_cache_loc(layer, forward_batch)
|
||||
|
||||
# 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(
|
||||
k=k,
|
||||
@@ -598,7 +598,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
),
|
||||
(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
|
||||
]
|
||||
else:
|
||||
@@ -611,7 +611,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
metadata.cu_seqlens_k = torch.nn.functional.pad(
|
||||
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
|
||||
]
|
||||
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),
|
||||
(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
|
||||
]
|
||||
|
||||
@@ -645,7 +645,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
metadata.cu_seqlens_k = torch.nn.functional.pad(
|
||||
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
|
||||
]
|
||||
|
||||
@@ -713,7 +713,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
else:
|
||||
# Use original set_kv_buffer path
|
||||
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
|
||||
)
|
||||
|
||||
@@ -721,7 +721,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
if self.data_type == torch.float8_e4m3fn and (not self.is_xqa_impl):
|
||||
q = q.to(torch.float8_e4m3fn)
|
||||
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:
|
||||
# [num_pages, page_size, num_kv_heads, head_dim] -> [num_pages, num_kv_heads, page_size, head_dim]
|
||||
k_cache = k_cache.view(
|
||||
@@ -799,7 +799,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
else:
|
||||
# Use original set_kv_buffer path
|
||||
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
|
||||
)
|
||||
|
||||
@@ -807,7 +807,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
q = q.to(torch.float8_e4m3fn)
|
||||
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]
|
||||
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(
|
||||
-1, self.page_size, layer.tp_k_head_num, layer.head_dim
|
||||
).permute(0, 2, 1, 3)
|
||||
|
||||
@@ -898,7 +898,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
assert (
|
||||
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."
|
||||
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
|
||||
)
|
||||
|
||||
@@ -924,7 +924,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
query = query.unsqueeze(1)
|
||||
|
||||
# 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)
|
||||
|
||||
# Get metadata
|
||||
@@ -1005,7 +1005,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
assert (
|
||||
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."
|
||||
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
|
||||
)
|
||||
|
||||
@@ -1046,7 +1046,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
# Ensure query has shape [bs, num_draft_tokens, num_q_heads, head_dim]
|
||||
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)
|
||||
|
||||
q = q.to(self.data_type)
|
||||
|
||||
@@ -117,6 +117,11 @@ class WaveAttnBackend(AttentionBackend):
|
||||
|
||||
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
|
||||
|
||||
if kv_indptr_buf is None:
|
||||
@@ -556,7 +561,7 @@ class WaveAttnBackend(AttentionBackend):
|
||||
o = torch.empty_like(q)
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
@@ -571,8 +576,8 @@ class WaveAttnBackend(AttentionBackend):
|
||||
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
|
||||
k.contiguous(),
|
||||
v.contiguous(),
|
||||
forward_batch.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_key_buffer(layer.layer_id),
|
||||
self.token_to_kv_pool.get_value_buffer(layer.layer_id),
|
||||
self.forward_metadata.qo_indptr,
|
||||
self.forward_metadata.kv_indptr,
|
||||
self.forward_metadata.kv_indices,
|
||||
@@ -606,14 +611,14 @@ class WaveAttnBackend(AttentionBackend):
|
||||
o = torch.empty_like(q)
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
self.decode_attention_fwd(
|
||||
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
|
||||
forward_batch.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_key_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),
|
||||
self.forward_metadata.kv_indptr,
|
||||
self.forward_metadata.kv_indices,
|
||||
|
||||
@@ -61,6 +61,10 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
self.device = model_runner.device
|
||||
self.decode_cuda_graph_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.kv_cache_dtype = model_runner.kv_cache_dtype
|
||||
self.kv_cache_dtype_str = model_runner.server_args.kv_cache_dtype
|
||||
@@ -122,7 +126,7 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
),
|
||||
(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
|
||||
]
|
||||
else:
|
||||
@@ -142,7 +146,7 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
),
|
||||
(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
|
||||
]
|
||||
|
||||
@@ -186,7 +190,7 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
metadata.cu_seqlens_k = torch.nn.functional.pad(
|
||||
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
|
||||
]
|
||||
# TODO: we need to test this part for llama 4 eagle case
|
||||
@@ -214,7 +218,7 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
),
|
||||
(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
|
||||
]
|
||||
|
||||
@@ -236,7 +240,7 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
),
|
||||
(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
|
||||
]
|
||||
|
||||
@@ -297,7 +301,7 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
)
|
||||
_, sort_order = torch.sort(keys, dim=1)
|
||||
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, :
|
||||
]
|
||||
.gather(1, cols)
|
||||
@@ -324,7 +328,7 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
metadata.cu_seqlens_k = torch.nn.functional.pad(
|
||||
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
|
||||
]
|
||||
|
||||
@@ -357,12 +361,12 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
(1, 0),
|
||||
)
|
||||
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
|
||||
]
|
||||
|
||||
# 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,
|
||||
metadata.encoder_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
|
||||
)
|
||||
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
|
||||
)
|
||||
else:
|
||||
forward_batch.token_to_kv_pool.set_mla_kv_buffer(
|
||||
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
k,
|
||||
@@ -501,9 +505,7 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
# Use Flash Attention for prefill
|
||||
if not self.use_mla:
|
||||
# Do multi-head attention
|
||||
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
|
||||
layer.layer_id
|
||||
)
|
||||
key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
||||
key_cache = key_cache.view(
|
||||
-1, self.page_size, layer.tp_k_head_num, layer.head_dim
|
||||
)
|
||||
@@ -614,9 +616,9 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
return output
|
||||
else:
|
||||
# Do absorbed multi-latent attention
|
||||
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer(
|
||||
layer.layer_id
|
||||
).to(q.dtype)
|
||||
kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
|
||||
q.dtype
|
||||
)
|
||||
k_rope = kv_cache[:, :, layer.v_head_dim :]
|
||||
c_kv = kv_cache[:, :, : layer.v_head_dim]
|
||||
k_rope_cache = k_rope.view(
|
||||
@@ -710,14 +712,14 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
else forward_batch.encoder_out_cache_loc
|
||||
)
|
||||
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
|
||||
)
|
||||
else:
|
||||
k_rope_val = (
|
||||
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,
|
||||
cache_loc,
|
||||
k,
|
||||
@@ -768,9 +770,7 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
if not self.use_mla:
|
||||
# Do multi-head attention
|
||||
|
||||
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
|
||||
layer.layer_id
|
||||
)
|
||||
key_cache, value_cache = self.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
||||
key_cache = key_cache.view(
|
||||
-1, self.page_size, layer.tp_k_head_num, layer.head_dim
|
||||
)
|
||||
@@ -876,9 +876,7 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
o = result
|
||||
else:
|
||||
# Do absorbed multi-latent attention
|
||||
kv_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id).to(
|
||||
q.dtype
|
||||
)
|
||||
kv_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(q.dtype)
|
||||
assert not use_cascade_attn, "Cascade attention is not supported with MLA"
|
||||
|
||||
if q_rope is not None:
|
||||
|
||||
@@ -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 (
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -135,7 +136,7 @@ class RadixAttention(nn.Module):
|
||||
)
|
||||
return output
|
||||
else:
|
||||
return forward_batch.attn_backend.forward(
|
||||
return get_attn_backend().forward(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
@@ -188,7 +189,7 @@ def unified_attention_with_output(
|
||||
# the FA kernel validates out.size(0) == q.size(0).
|
||||
forward_batch._attn_output = output[:real_num_tokens]
|
||||
|
||||
ret = forward_batch.attn_backend.forward(
|
||||
ret = get_attn_backend().forward(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
|
||||
@@ -22,6 +22,7 @@ from torch import nn
|
||||
|
||||
from sglang.srt.compilation.compilation_config import register_split_op
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -92,7 +93,7 @@ class RadixLinearAttention(nn.Module):
|
||||
)
|
||||
return output
|
||||
else:
|
||||
return forward_batch.attn_backend.forward(
|
||||
return get_attn_backend().forward(
|
||||
layer=self,
|
||||
forward_batch=forward_batch,
|
||||
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.
|
||||
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,
|
||||
forward_batch=forward_batch,
|
||||
mixed_qkv=mixed_qkv[:real_num_tokens],
|
||||
|
||||
@@ -14,6 +14,7 @@ from sglang.srt.layers.dp_attention import (
|
||||
get_attention_cp_size,
|
||||
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
|
||||
|
||||
|
||||
@@ -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()
|
||||
)
|
||||
|
||||
forward_batch.token_to_kv_pool.set_kv_buffer(
|
||||
get_token_to_kv_pool().set_kv_buffer(
|
||||
layer,
|
||||
cache_loc,
|
||||
key_cache_full,
|
||||
|
||||
@@ -57,6 +57,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
CaptureHiddenMode,
|
||||
PPProxyTensors,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
||||
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
|
||||
PiecewiseCudaGraphRunner,
|
||||
freeze_gc,
|
||||
@@ -292,9 +293,6 @@ class BreakableCudaGraphRunner:
|
||||
next_token_logits_buffer=None,
|
||||
orig_seq_lens=orig_seq_lens,
|
||||
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],
|
||||
seq_lens_sum=num_tokens,
|
||||
mamba_track_indices=None,
|
||||
@@ -329,8 +327,11 @@ class BreakableCudaGraphRunner:
|
||||
"""Warmup the model with a forward pass."""
|
||||
num_tokens = self.capture_num_tokens[0]
|
||||
forward_batch = self._build_capture_forward_batch(num_tokens)
|
||||
self.model_runner.attn_backend.init_forward_metadata(forward_batch)
|
||||
self._run_forward(forward_batch, num_tokens)
|
||||
with forward_context(
|
||||
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):
|
||||
"""Capture breakable CUDA graphs for all token sizes."""
|
||||
@@ -394,14 +395,17 @@ class BreakableCudaGraphRunner:
|
||||
self.model_runner.token_to_kv_pool.invalidate_loc_cache()
|
||||
return self._run_forward(forward_batch, num_tokens)
|
||||
|
||||
for _ in range(2):
|
||||
self.device_module.synchronize()
|
||||
self.model_runner.tp_group.barrier()
|
||||
run_once()
|
||||
with forward_context(
|
||||
ForwardContext(attn_backend=self.model_runner.attn_backend)
|
||||
):
|
||||
for _ in range(2):
|
||||
self.device_module.synchronize()
|
||||
self.model_runner.tp_group.barrier()
|
||||
run_once()
|
||||
|
||||
graph = BreakableCUDAGraph()
|
||||
with BreakableCUDAGraphCapture(cuda_graph=graph, pool=pool, stream=stream):
|
||||
output = run_once()
|
||||
graph = BreakableCUDAGraph()
|
||||
with BreakableCUDAGraphCapture(cuda_graph=graph, pool=pool, stream=stream):
|
||||
output = run_once()
|
||||
|
||||
return graph, output
|
||||
|
||||
|
||||
@@ -36,6 +36,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
PPProxyTensors,
|
||||
enable_num_token_non_padded,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
||||
from sglang.srt.utils import (
|
||||
log_info_on_rank0,
|
||||
require_attn_tp_gather,
|
||||
@@ -679,9 +680,6 @@ class CPUGraphRunner:
|
||||
input_ids=input_ids,
|
||||
req_pool_indices=req_pool_indices,
|
||||
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,
|
||||
seq_lens_sum=seq_lens.sum().item(),
|
||||
return_logprob=False,
|
||||
@@ -693,43 +691,46 @@ class CPUGraphRunner:
|
||||
num_token_non_padded=self.num_token_non_padded,
|
||||
global_forward_mode=self.capture_forward_mode,
|
||||
)
|
||||
self.model_runner.attn_backend.init_forward_metadata_capture_cpu_graph(
|
||||
bs,
|
||||
num_tokens,
|
||||
req_pool_indices,
|
||||
seq_lens,
|
||||
None,
|
||||
forward_batch.forward_mode,
|
||||
forward_batch.spec_info,
|
||||
)
|
||||
# Do infernence to avoid setting attr at runtime, e.g.,
|
||||
# self.attn_mha.kv_b_proj = self.kv_b_proj for full graph compile on CPU
|
||||
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 forward_context(
|
||||
ForwardContext(attn_backend=self.model_runner.attn_backend)
|
||||
):
|
||||
self.model_runner.attn_backend.init_forward_metadata_capture_cpu_graph(
|
||||
bs,
|
||||
num_tokens,
|
||||
req_pool_indices,
|
||||
seq_lens,
|
||||
None,
|
||||
forward_batch.forward_mode,
|
||||
forward_batch.spec_info,
|
||||
)
|
||||
|
||||
# Run and capture
|
||||
def run_once():
|
||||
# Clean intermediate result cache for DP attention
|
||||
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
|
||||
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):
|
||||
with torch.no_grad():
|
||||
self.model_runner.tp_group.barrier()
|
||||
out = run_once()
|
||||
# Save the captured forward_batch
|
||||
self.captured_forward_batches[bs] = forward_batch
|
||||
return forward, out
|
||||
self.model_runner.model.forward(
|
||||
forward_batch.input_ids,
|
||||
forward_batch.positions,
|
||||
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):
|
||||
|
||||
|
||||
@@ -65,6 +65,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
compute_local_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.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups
|
||||
from sglang.srt.utils import (
|
||||
@@ -1016,9 +1017,6 @@ class CudaGraphRunner:
|
||||
seq_lens_cpu=seq_lens_cpu,
|
||||
next_token_logits_buffer=next_token_logits_buffer,
|
||||
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,
|
||||
seq_lens_sum=seq_lens.sum().item(),
|
||||
mamba_track_indices=mamba_track_indices,
|
||||
@@ -1040,85 +1038,90 @@ class CudaGraphRunner:
|
||||
lora_ids=lora_ids,
|
||||
)
|
||||
|
||||
# HiSparse: set coordinator so the hisparse code path is captured into the graph
|
||||
forward_batch.hisparse_coordinator = self.model_runner.hisparse_coordinator
|
||||
if forward_batch.hisparse_coordinator is not None:
|
||||
forward_batch.hisparse_coordinator.num_real_reqs.fill_(bs)
|
||||
# Trip the coordinator so the hisparse code path is captured into the
|
||||
# graph; backends read it from self.model_runner.hisparse_coordinator.
|
||||
hisparse_coordinator = self.model_runner.hisparse_coordinator
|
||||
if hisparse_coordinator is not None:
|
||||
hisparse_coordinator.num_real_reqs.fill_(bs)
|
||||
|
||||
if buffers.ngram_embedding_info is not None:
|
||||
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:
|
||||
self.model_runner.lora_manager.prepare_lora_batch(forward_batch)
|
||||
if lora_ids is not None:
|
||||
self.model_runner.lora_manager.prepare_lora_batch(forward_batch)
|
||||
|
||||
# Attention backend
|
||||
attn_backend.init_forward_metadata_capture_cuda_graph(
|
||||
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,
|
||||
attn_backend.init_forward_metadata_capture_cuda_graph(
|
||||
bs,
|
||||
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 = {}
|
||||
if (
|
||||
self.pp_size > 1
|
||||
and "pp_proxy_tensors" in inspect.signature(forward).parameters
|
||||
):
|
||||
kwargs["pp_proxy_tensors"] = PPProxyTensors(
|
||||
{k: v.clone() for k, v in pp_proxy_tensors.tensors.items()}
|
||||
def run_once():
|
||||
# Without this, warmup-1 caches the translation; the capture
|
||||
# run hits the cache, 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()
|
||||
|
||||
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = (
|
||||
None
|
||||
)
|
||||
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]
|
||||
set_dp_buffer_len(
|
||||
global_dp_buffer_len,
|
||||
num_tokens,
|
||||
forward_batch.dp_padding_mode.is_max_len(),
|
||||
)
|
||||
set_is_extend_in_batch(False)
|
||||
|
||||
logits_output_or_pp_proxy_tensors = forward(
|
||||
input_ids,
|
||||
forward_batch.positions,
|
||||
forward_batch,
|
||||
**kwargs,
|
||||
kwargs = {}
|
||||
if (
|
||||
self.pp_size > 1
|
||||
and "pp_proxy_tensors" in inspect.signature(forward).parameters
|
||||
):
|
||||
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
|
||||
|
||||
|
||||
@@ -9,6 +9,10 @@ import triton.language as tl
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
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:
|
||||
@@ -55,6 +59,7 @@ class ForwardBatchDeepSeekMHAMixin:
|
||||
|
||||
def prepare_chunked_kv_indices(self, device: torch.device):
|
||||
self.prefix_chunk_kv_indices = []
|
||||
req_to_token = get_req_to_token_pool().req_to_token
|
||||
for idx in range(self.num_prefix_chunks):
|
||||
chunk_starts = self.prefix_chunk_starts[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,)](
|
||||
self.req_to_token_pool.req_to_token,
|
||||
req_to_token,
|
||||
self.req_pool_indices,
|
||||
chunk_starts,
|
||||
chunk_seq_lens,
|
||||
chunk_cu_seq_lens,
|
||||
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)
|
||||
|
||||
@@ -111,7 +116,7 @@ class ForwardBatchDeepSeekMHAMixin:
|
||||
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool
|
||||
|
||||
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"
|
||||
|
||||
if not any(self.extend_prefix_lens_cpu):
|
||||
@@ -191,14 +196,15 @@ class ForwardBatchDeepSeekMHAMixin:
|
||||
device=self.req_pool_indices.device,
|
||||
)
|
||||
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,)](
|
||||
self.req_to_token_pool.req_to_token,
|
||||
req_to_token,
|
||||
self.req_pool_indices,
|
||||
self.seq_lens,
|
||||
kv_indptr,
|
||||
None,
|
||||
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
|
||||
return kv_indices
|
||||
|
||||
@@ -63,11 +63,8 @@ from sglang.srt.utils import (
|
||||
from sglang.srt.utils.common import ceil_align
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
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.mem_cache.memory_pool import KVCache, ReqToTokenPool
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
from sglang.srt.speculative.spec_info import SpecInput, SpeculativeAlgorithm
|
||||
@@ -369,11 +366,6 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
# Sampling info
|
||||
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
|
||||
original_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)
|
||||
return_pooled_hidden_states: bool = False
|
||||
|
||||
# For hisparse
|
||||
hisparse_coordinator: Optional[HiSparseCoordinator] = None
|
||||
|
||||
# For ngram embedding
|
||||
ngram_embedding_info: Optional[NgramEmbeddingInfo] = None
|
||||
|
||||
@@ -536,9 +525,6 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
multi_item_delimiter_indices=batch.multi_item_delimiter_indices,
|
||||
lora_ids=[req.lora_id for req in batch.reqs],
|
||||
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_info=batch.spec_info,
|
||||
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)
|
||||
@@ -146,6 +146,11 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardMode,
|
||||
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.model_runner_kv_cache_mixin import (
|
||||
ModelRunnerKVCacheMixin,
|
||||
@@ -2638,9 +2643,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
seq_lens_cpu=buffers.seq_lens_cpu,
|
||||
next_token_logits_buffer=buffers.next_token_logits_buffer,
|
||||
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,
|
||||
seq_lens_sum=buffers.seq_lens.sum().item(),
|
||||
encoder_lens=buffers.encoder_lens,
|
||||
@@ -2701,8 +2703,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
|
||||
torch.get_device_module(self.device).synchronize()
|
||||
self.tp_group.barrier()
|
||||
with torch.inference_mode(), run_ctx or empty_context():
|
||||
run_once()
|
||||
with forward_context(ForwardContext(attn_backend=self.attn_backend)):
|
||||
with torch.inference_mode(), run_ctx or empty_context():
|
||||
run_once()
|
||||
|
||||
def maybe_init_ngram_embedding(self):
|
||||
self.use_ngram_embedding = self.model_config.use_ngram_embedding
|
||||
@@ -2979,6 +2982,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
pp_proxy_tensors=None,
|
||||
) -> Union[LogitsProcessorOutput, PPProxyTensors]:
|
||||
# Set extra arguments
|
||||
pdmux_override = False
|
||||
if not skip_attn_backend_init:
|
||||
if hasattr(self.model, "prepare_forward_batch"):
|
||||
# Prepare model-specific attention metadata before planning,
|
||||
@@ -2986,7 +2990,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
self.model.prepare_forward_batch(forward_batch)
|
||||
if self.server_args.enable_pdmux:
|
||||
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:
|
||||
self.attn_backend.init_forward_metadata(forward_batch)
|
||||
# FIXME: add pp_proxy_tensors arg to all models
|
||||
@@ -3000,7 +3007,8 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
if self.device_timer
|
||||
else contextlib.nullcontext()
|
||||
)
|
||||
with ctx:
|
||||
|
||||
def _do_forward():
|
||||
return self.model.forward(
|
||||
forward_batch.input_ids,
|
||||
forward_batch.positions,
|
||||
@@ -3008,6 +3016,14 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
**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(
|
||||
self,
|
||||
forward_batch: ForwardBatch,
|
||||
@@ -3216,93 +3232,100 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
reinit_attn_backend: bool = False,
|
||||
split_forward_count: int = 1,
|
||||
) -> ModelRunnerOutput:
|
||||
# Check whether can run cuda graph
|
||||
mode_check = (
|
||||
forward_batch.forward_mode.is_cpu_graph
|
||||
if self.device == "cpu"
|
||||
else forward_batch.forward_mode.is_cuda_graph
|
||||
)
|
||||
can_run_graph = bool(
|
||||
mode_check()
|
||||
and self.graph_runner
|
||||
and self.graph_runner.can_run(forward_batch)
|
||||
)
|
||||
|
||||
# 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,
|
||||
# Honor an outer-published context (spec workers wrap each per-step
|
||||
# draft forward with the i-th child backend); otherwise publish this
|
||||
# runner's own attn_backend for the forward.
|
||||
if has_forward_context():
|
||||
ctx_mgr = contextlib.nullcontext()
|
||||
else:
|
||||
ctx_mgr = forward_context(ForwardContext(attn_backend=self.attn_backend))
|
||||
with ctx_mgr:
|
||||
mode_check = (
|
||||
forward_batch.forward_mode.is_cpu_graph
|
||||
if self.device == "cpu"
|
||||
else forward_batch.forward_mode.is_cuda_graph
|
||||
)
|
||||
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)
|
||||
|
||||
# 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(
|
||||
self, logits_output: LogitsProcessorOutput, sampling_info: SamplingBatchInfo
|
||||
):
|
||||
|
||||
@@ -58,6 +58,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardMode,
|
||||
PPProxyTensors,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
||||
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
||||
from sglang.srt.utils import (
|
||||
get_available_gpu_memory,
|
||||
@@ -387,9 +388,6 @@ class PiecewiseCudaGraphRunner:
|
||||
next_token_logits_buffer=None,
|
||||
orig_seq_lens=torch.tensor([num_tokens], device=self.device),
|
||||
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,
|
||||
seq_lens_sum=num_tokens,
|
||||
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
|
||||
set_dp_buffer_len(None, num_tokens, forward_batch.dp_padding_mode.is_max_len())
|
||||
set_is_extend_in_batch(False)
|
||||
with set_forward_context(
|
||||
forward_batch,
|
||||
self.attention_layers,
|
||||
self.quant_config,
|
||||
self.moe_layers,
|
||||
self.moe_fusions,
|
||||
with forward_context(
|
||||
ForwardContext(attn_backend=self.model_runner.attn_backend)
|
||||
):
|
||||
_ = self.model_runner.model.forward(
|
||||
forward_batch.input_ids,
|
||||
forward_batch.positions,
|
||||
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,
|
||||
)
|
||||
|
||||
def _cache_loc_dtype(self):
|
||||
return torch.int64 if not is_npu() else torch.int32
|
||||
@@ -554,9 +555,6 @@ class PiecewiseCudaGraphRunner:
|
||||
next_token_logits_buffer=None,
|
||||
orig_seq_lens=torch.tensor([num_tokens], device=self.device),
|
||||
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,
|
||||
seq_lens_sum=num_tokens,
|
||||
mamba_track_indices=mamba_track_indices,
|
||||
@@ -586,52 +584,59 @@ class PiecewiseCudaGraphRunner:
|
||||
lora_ids=None,
|
||||
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)
|
||||
|
||||
if lora_ids is not None:
|
||||
self.model_runner.lora_manager.prepare_lora_batch(forward_batch)
|
||||
if lora_ids is not None:
|
||||
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
|
||||
def run_once():
|
||||
# Invalidate SWA loc cache — same fix as in cuda_graph_runner.run_once.
|
||||
if self.model_runner.is_hybrid_swa:
|
||||
self.model_runner.token_to_kv_pool.invalidate_loc_cache()
|
||||
# Run and capture
|
||||
def run_once():
|
||||
# Invalidate SWA loc cache — same fix as in cuda_graph_runner.run_once.
|
||||
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,
|
||||
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,
|
||||
# Clean intermediate result cache for DP attention
|
||||
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = (
|
||||
None
|
||||
)
|
||||
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
|
||||
# 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()
|
||||
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
|
||||
|
||||
# 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
|
||||
|
||||
@@ -733,9 +738,6 @@ class PiecewiseCudaGraphRunner:
|
||||
next_token_logits_buffer=next_token_logits_buffer,
|
||||
orig_seq_lens=forward_batch.orig_seq_lens,
|
||||
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,
|
||||
seq_lens_sum=forward_batch.seq_lens_sum,
|
||||
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.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 (
|
||||
AttnForwardMethod,
|
||||
)
|
||||
@@ -153,7 +154,7 @@ def handle_attention_dsa(attn, forward_batch):
|
||||
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
|
||||
backend = backend.primary
|
||||
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.communicator import get_attn_tp_context
|
||||
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 (
|
||||
_is_cuda,
|
||||
_is_hip,
|
||||
@@ -38,7 +42,7 @@ if _use_aiter_gfx95:
|
||||
|
||||
|
||||
def _resolve_attn_backend(forward_batch: ForwardBatch):
|
||||
backend = forward_batch.attn_backend
|
||||
backend = get_attn_backend()
|
||||
if isinstance(backend, TboAttnBackend):
|
||||
backend = backend.primary
|
||||
return backend
|
||||
@@ -334,8 +338,8 @@ class DeepseekMHAForwardMixin:
|
||||
# Only initialize the info once
|
||||
if has_extend_prefix and forward_batch.num_prefix_chunks is None:
|
||||
forward_batch.prepare_chunked_prefix_cache_info(q.device)
|
||||
if hasattr(forward_batch.attn_backend, "init_mha_chunk_metadata"):
|
||||
forward_batch.attn_backend.init_mha_chunk_metadata(forward_batch)
|
||||
if hasattr(get_attn_backend(), "init_mha_chunk_metadata"):
|
||||
get_attn_backend().init_mha_chunk_metadata(forward_batch)
|
||||
|
||||
forward_batch.mha_return_lse = has_extend_prefix
|
||||
# Do mha for extended part without prefix
|
||||
@@ -380,8 +384,8 @@ class DeepseekMHAForwardMixin:
|
||||
# Only initialize the info once
|
||||
if has_extend_prefix and forward_batch.num_prefix_chunks is None:
|
||||
forward_batch.num_prefix_chunks = 0
|
||||
if hasattr(forward_batch.attn_backend, "init_mha_chunk_metadata"):
|
||||
forward_batch.attn_backend.init_mha_chunk_metadata(forward_batch)
|
||||
if hasattr(get_attn_backend(), "init_mha_chunk_metadata"):
|
||||
get_attn_backend().init_mha_chunk_metadata(forward_batch)
|
||||
forward_batch.mha_return_lse = False
|
||||
# Do mha for extended part without prefix
|
||||
forward_batch.set_attn_attend_prefix_cache(False)
|
||||
@@ -449,12 +453,12 @@ class DeepseekMHAForwardMixin:
|
||||
):
|
||||
if _is_cuda or _use_aiter_gfx95:
|
||||
# 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
|
||||
)
|
||||
elif _is_npu:
|
||||
# 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
|
||||
)
|
||||
else:
|
||||
@@ -462,7 +466,7 @@ class DeepseekMHAForwardMixin:
|
||||
latent_cache[:, :, self.kv_lora_rank :] = k_pe.clone()
|
||||
|
||||
# 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
|
||||
)
|
||||
|
||||
@@ -473,12 +477,12 @@ class DeepseekMHAForwardMixin:
|
||||
forward_batch: ForwardBatch,
|
||||
):
|
||||
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
|
||||
)
|
||||
kv_a = kv_a.squeeze(1)
|
||||
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
|
||||
)
|
||||
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
|
||||
"""
|
||||
backend = forward_batch.attn_backend
|
||||
backend = get_attn_backend()
|
||||
if isinstance(backend, TboAttnBackend): # if enable tbo, get primary backend
|
||||
backend = backend.primary
|
||||
kv_indices = backend.forward_metadata.page_table_1_flattened
|
||||
@@ -506,9 +510,7 @@ class DeepseekMHAForwardMixin:
|
||||
kv_indices is not None
|
||||
), "page_table_1_flattened should have been generated for FP8 MHA path"
|
||||
|
||||
kv_cache_fp8 = forward_batch.token_to_kv_pool.get_key_buffer(
|
||||
self.attn_mha.layer_id
|
||||
)
|
||||
kv_cache_fp8 = get_token_to_kv_pool().get_key_buffer(self.attn_mha.layer_id)
|
||||
|
||||
kv_latent_bf16 = dequantize_k_cache_paged(kv_cache_fp8, kv_indices)
|
||||
|
||||
@@ -544,7 +546,7 @@ class DeepseekMHAForwardMixin:
|
||||
self.current_attention_backend == "fa3"
|
||||
and self.kv_cache_dtype != "auto"
|
||||
):
|
||||
attn_dtype = forward_batch.token_to_kv_pool.dtype
|
||||
attn_dtype = get_token_to_kv_pool().dtype
|
||||
else:
|
||||
attn_dtype = k_nope.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,
|
||||
)
|
||||
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 (
|
||||
FORWARD_ABSORB_CORE_ATTENTION_BACKENDS,
|
||||
_is_cpu,
|
||||
@@ -420,9 +424,7 @@ class DeepseekMLAForwardMixin:
|
||||
q_pe,
|
||||
k_nope,
|
||||
k_pe,
|
||||
forward_batch.token_to_kv_pool.get_key_buffer(
|
||||
self.attn_mqa.layer_id
|
||||
),
|
||||
get_token_to_kv_pool().get_key_buffer(self.attn_mqa.layer_id),
|
||||
forward_batch.out_cache_loc,
|
||||
positions,
|
||||
cos,
|
||||
@@ -516,9 +518,7 @@ class DeepseekMLAForwardMixin:
|
||||
q_pe,
|
||||
k_nope,
|
||||
k_pe,
|
||||
forward_batch.token_to_kv_pool.get_key_buffer(
|
||||
self.attn_mqa.layer_id
|
||||
),
|
||||
get_token_to_kv_pool().get_key_buffer(self.attn_mqa.layer_id),
|
||||
forward_batch.out_cache_loc,
|
||||
positions,
|
||||
cos,
|
||||
@@ -694,7 +694,7 @@ class DeepseekMLAForwardMixin:
|
||||
return (
|
||||
get_global_server_args().dsa_decode_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 (
|
||||
self.current_attention_backend in ("trtllm_mla", "tokenspeed_mla")
|
||||
@@ -702,7 +702,7 @@ class DeepseekMLAForwardMixin:
|
||||
forward_batch.forward_mode.is_decode_or_idle()
|
||||
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:
|
||||
|
||||
+9
-7
@@ -7,6 +7,10 @@ import torch
|
||||
|
||||
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_context import (
|
||||
get_attn_backend,
|
||||
get_token_to_kv_pool,
|
||||
)
|
||||
from sglang.srt.models.deepseek_common.utils import (
|
||||
_is_cuda,
|
||||
_is_hip,
|
||||
@@ -108,10 +112,10 @@ class DeepseekMLARocmForwardMixin:
|
||||
device=q.device,
|
||||
)
|
||||
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
|
||||
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
|
||||
if attn_logits is None:
|
||||
attn_logits = torch.empty(
|
||||
@@ -126,12 +130,10 @@ class DeepseekMLARocmForwardMixin:
|
||||
)
|
||||
|
||||
# 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
|
||||
)
|
||||
key_cache_buf = forward_batch.token_to_kv_pool.get_key_buffer(
|
||||
self.attn_mqa.layer_id
|
||||
)
|
||||
key_cache_buf = get_token_to_kv_pool().get_key_buffer(self.attn_mqa.layer_id)
|
||||
val_cache_buf = key_cache_buf[..., : self.kv_lora_rank]
|
||||
|
||||
return (
|
||||
@@ -194,7 +196,7 @@ class DeepseekMLARocmForwardMixin:
|
||||
|
||||
if enable_rope_fusion:
|
||||
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
|
||||
)
|
||||
|
||||
|
||||
@@ -78,6 +78,10 @@ from sglang.srt.model_executor.cuda_graph_runner import (
|
||||
get_is_capture_mode,
|
||||
)
|
||||
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.weight_utils import default_weight_loader
|
||||
from sglang.srt.models.dbrx import ReplicatedLinear
|
||||
@@ -398,7 +402,7 @@ class MQALayer(nn.Module):
|
||||
kv = qkv_a[..., self.q_lora_rank :]
|
||||
else:
|
||||
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:
|
||||
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
|
||||
token_to_kv_pool.set_swa_key_buffer_radix_fused_norm_rope(
|
||||
@@ -467,6 +471,7 @@ class MQALayer(nn.Module):
|
||||
x=x,
|
||||
q_lora=q_lora,
|
||||
forward_batch=forward_batch,
|
||||
attn_backend=attn_backend,
|
||||
enable_multi_stream=True,
|
||||
q_lora_ready=q_lora_ready,
|
||||
)
|
||||
@@ -533,7 +538,12 @@ class MQALayer(nn.Module):
|
||||
del qkv_a
|
||||
|
||||
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:
|
||||
attn_backend.forward_core_compressor(
|
||||
x,
|
||||
@@ -556,7 +566,7 @@ class MQALayer(nn.Module):
|
||||
), "short-circuiting allreduce will lead to hangs"
|
||||
return x
|
||||
|
||||
attn_backend = forward_batch.attn_backend
|
||||
attn_backend = get_attn_backend()
|
||||
if TYPE_CHECKING:
|
||||
assert isinstance(
|
||||
attn_backend,
|
||||
@@ -1130,7 +1140,7 @@ class DeepseekV4Model(nn.Module):
|
||||
|
||||
# Upgrade lazy raw metadata on the main stream once before any layer
|
||||
# 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):
|
||||
layer = self.layers[i]
|
||||
@@ -1278,15 +1288,14 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
forward_batch.seq_lens_cpu.tolist(),
|
||||
)
|
||||
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.apply_cp_reindex()
|
||||
core_meta.init_flashmla_related()
|
||||
if metadata.indexer_metadata is not None:
|
||||
metadata.indexer_metadata = (
|
||||
forward_batch.attn_backend.init_forward_metadata_indexer(
|
||||
core_meta
|
||||
)
|
||||
attn_backend.init_forward_metadata_indexer(core_meta)
|
||||
)
|
||||
|
||||
with get_attn_tp_context().maybe_input_scattered(forward_batch):
|
||||
|
||||
@@ -37,6 +37,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
||||
VocabParallelEmbedding,
|
||||
)
|
||||
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.server_args import get_global_server_args
|
||||
from sglang.srt.utils import add_prefix
|
||||
@@ -250,15 +251,14 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM):
|
||||
forward_batch.seq_lens_cpu.tolist(),
|
||||
)
|
||||
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.apply_cp_reindex()
|
||||
core_meta.init_flashmla_related()
|
||||
if metadata.indexer_metadata is not None:
|
||||
metadata.indexer_metadata = (
|
||||
forward_batch.attn_backend.init_forward_metadata_indexer(
|
||||
core_meta
|
||||
)
|
||||
attn_backend.init_forward_metadata_indexer(core_meta)
|
||||
)
|
||||
|
||||
hidden_states, pre_hc_head = self.model(input_ids, positions, forward_batch)
|
||||
|
||||
@@ -33,6 +33,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
||||
VocabParallelEmbedding,
|
||||
)
|
||||
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.server_args import get_global_server_args
|
||||
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
|
||||
|
||||
attn_backend = forward_batch.attn_backend
|
||||
attn_backend = get_attn_backend()
|
||||
assert isinstance(attn_backend, HybridLinearAttnBackend)
|
||||
assert isinstance(attn_backend.linear_attn_backend, Mamba2AttnBackend)
|
||||
# Mamba block
|
||||
|
||||
@@ -40,6 +40,7 @@ from sglang.srt.managers.schedule_batch import (
|
||||
flatten_nested_list,
|
||||
)
|
||||
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 (
|
||||
default_weight_loader,
|
||||
maybe_remap_kv_scale_name,
|
||||
@@ -220,7 +221,7 @@ class Gemma3ForConditionalGeneration(PreTrainedModel):
|
||||
mask_dtype: torch.dtype,
|
||||
):
|
||||
"""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
|
||||
bidirectional_attn_masks_list = []
|
||||
bidirectional_attn_mask_indptr = torch.zeros(
|
||||
@@ -265,10 +266,10 @@ class Gemma3ForConditionalGeneration(PreTrainedModel):
|
||||
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
|
||||
)
|
||||
forward_batch.attn_backend.forward_metadata.custom_mask = (
|
||||
get_attn_backend().forward_metadata.custom_mask = (
|
||||
bidirectional_attn_masks
|
||||
)
|
||||
|
||||
|
||||
@@ -52,6 +52,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardMode,
|
||||
PPProxyTensors,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_context import get_attn_backend
|
||||
from sglang.srt.model_loader.weight_utils import (
|
||||
default_weight_loader,
|
||||
maybe_remap_kv_scale_name,
|
||||
@@ -315,7 +316,7 @@ class Gemma4ForConditionalGeneration(PreTrainedModel):
|
||||
|
||||
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(
|
||||
"Bidirectional attention for image tokens requires TritonAttnBackend. "
|
||||
"Falling back to causal attention, which may degrade image quality."
|
||||
@@ -389,12 +390,10 @@ class Gemma4ForConditionalGeneration(PreTrainedModel):
|
||||
)
|
||||
if bidirectional_attn_masks_list:
|
||||
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
|
||||
)
|
||||
forward_batch.attn_backend.forward_metadata.custom_mask = (
|
||||
bidirectional_attn_masks
|
||||
)
|
||||
get_attn_backend().forward_metadata.custom_mask = bidirectional_attn_masks
|
||||
|
||||
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
|
||||
vt = self.vision_tower
|
||||
|
||||
@@ -29,6 +29,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
||||
VocabParallelEmbedding,
|
||||
)
|
||||
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.models.transformers import maybe_prefix
|
||||
from sglang.srt.utils import make_layers
|
||||
@@ -139,7 +140,7 @@ class GraniteMoeHybridMambaDecoderLayer(nn.Module):
|
||||
hidden_states = self.input_layernorm(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.linear_attn_backend, Mamba2AttnBackend)
|
||||
attn_backend.linear_attn_backend.forward(
|
||||
|
||||
@@ -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.vocab_parallel_embedding import ParallelLMHead
|
||||
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.models.qwen2 import Qwen2MLP, Qwen2Model
|
||||
from sglang.srt.utils import add_prefix
|
||||
@@ -258,11 +259,9 @@ class JetBlock(nn.Module):
|
||||
hidden_states: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
) -> torch.Tensor:
|
||||
assert isinstance(forward_batch.attn_backend, HybridLinearAttnBackend)
|
||||
assert isinstance(
|
||||
forward_batch.attn_backend.linear_attn_backend, MambaAttnBackendBase
|
||||
)
|
||||
linear_attn_backend = forward_batch.attn_backend.linear_attn_backend
|
||||
assert isinstance(get_attn_backend(), HybridLinearAttnBackend)
|
||||
assert isinstance(get_attn_backend().linear_attn_backend, MambaAttnBackendBase)
|
||||
linear_attn_backend = get_attn_backend().linear_attn_backend
|
||||
forward_metadata = linear_attn_backend.forward_metadata
|
||||
layer_cache = linear_attn_backend.req_to_token_pool.mamba2_layer_cache(
|
||||
self.layer_id
|
||||
|
||||
@@ -40,6 +40,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
||||
VocabParallelEmbedding,
|
||||
)
|
||||
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 (
|
||||
default_weight_loader,
|
||||
sharded_weight_loader,
|
||||
@@ -263,12 +264,10 @@ class Lfm2ShortConv(nn.Module):
|
||||
if forward_batch.forward_mode.is_idle():
|
||||
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]
|
||||
req_pool_indices = forward_batch.req_pool_indices
|
||||
mamba_indices = forward_batch.req_to_token_pool.get_mamba_indices(
|
||||
req_pool_indices
|
||||
)
|
||||
mamba_indices = get_req_to_token_pool().get_mamba_indices(req_pool_indices)
|
||||
|
||||
# Project and split into gates: B (pre-conv), C (post-conv), x (input)
|
||||
proj, _ = self.in_proj(hidden_states)
|
||||
|
||||
@@ -42,6 +42,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
||||
VocabParallelEmbedding,
|
||||
)
|
||||
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 (
|
||||
default_weight_loader,
|
||||
sharded_weight_loader,
|
||||
@@ -326,12 +327,10 @@ class Lfm2MoeShortConv(nn.Module):
|
||||
if forward_batch.forward_mode.is_idle():
|
||||
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]
|
||||
req_pool_indices = forward_batch.req_pool_indices
|
||||
mamba_indices = forward_batch.req_to_token_pool.get_mamba_indices(
|
||||
req_pool_indices
|
||||
)
|
||||
mamba_indices = get_req_to_token_pool().get_mamba_indices(req_pool_indices)
|
||||
|
||||
proj, _ = self.in_proj(hidden_states)
|
||||
B_gate, C_gate, x = proj.chunk(3, dim=-1)
|
||||
|
||||
@@ -14,6 +14,10 @@ from sglang.srt.distributed import (
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
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_context import (
|
||||
get_req_to_token_pool,
|
||||
get_token_to_kv_pool,
|
||||
)
|
||||
from sglang.srt.models.registry import import_model_classes
|
||||
from sglang.srt.utils import is_npu
|
||||
|
||||
@@ -221,9 +225,9 @@ class MindSporeForCausalLM(torch.nn.Module):
|
||||
def prepare_cache(cache_list, is_key_cache):
|
||||
for i in range(self.config.num_hidden_layers):
|
||||
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:
|
||||
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)
|
||||
if self.use_mla and cache_ms.ndim == 3:
|
||||
cache_ms = mint.unsqueeze(cache_ms, 2)
|
||||
@@ -275,10 +279,10 @@ class MindSporeForCausalLM(torch.nn.Module):
|
||||
if forward_batch.forward_mode.is_target_verify():
|
||||
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(
|
||||
(
|
||||
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()
|
||||
][:, ::page_size]
|
||||
// page_size
|
||||
|
||||
@@ -69,6 +69,7 @@ from sglang.srt.model_executor.breakable_cuda_graph.context import (
|
||||
is_in_breakable_cuda_graph,
|
||||
)
|
||||
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,
|
||||
maybe_remap_kv_scale_name,
|
||||
@@ -414,7 +415,7 @@ class NemotronHMambaDecoderLayer(nn.Module):
|
||||
) -> torch.Tensor:
|
||||
"""Core Mamba forward logic, called directly or via split op."""
|
||||
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.linear_attn_backend, Mamba2AttnBackend)
|
||||
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
|
||||
# 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
|
||||
num_actual_tokens = metadata.num_prefill_tokens + (
|
||||
metadata.num_decodes * metadata.draft_token_num
|
||||
|
||||
@@ -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.vocab_parallel_embedding import ParallelLMHead
|
||||
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 (
|
||||
default_weight_loader,
|
||||
maybe_remap_kv_scale_name,
|
||||
@@ -214,7 +215,7 @@ class Qwen3Attention(nn.Module):
|
||||
|
||||
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)
|
||||
slot_mapping = forward_batch.out_cache_loc
|
||||
|
||||
|
||||
@@ -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.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.models.bailing_moe import BailingMoEForCausalLM
|
||||
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"
|
||||
and self.kv_cache_dtype != "auto"
|
||||
):
|
||||
attn_dtype = forward_batch.token_to_kv_pool.dtype
|
||||
attn_dtype = get_token_to_kv_pool().dtype
|
||||
else:
|
||||
attn_dtype = k_nope.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[..., 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,
|
||||
forward_batch.out_cache_loc,
|
||||
k_nope,
|
||||
@@ -701,8 +705,8 @@ class SarvamMoEMLAAttention(nn.Module):
|
||||
forward_batch.prepare_chunked_prefix_cache_info(q.device)
|
||||
else:
|
||||
forward_batch.num_prefix_chunks = 0
|
||||
if hasattr(forward_batch.attn_backend, "init_mha_chunk_metadata"):
|
||||
forward_batch.attn_backend.init_mha_chunk_metadata(forward_batch)
|
||||
if hasattr(get_attn_backend(), "init_mha_chunk_metadata"):
|
||||
get_attn_backend().init_mha_chunk_metadata(forward_batch)
|
||||
|
||||
forward_batch.set_attn_attend_prefix_cache(False)
|
||||
forward_batch.mha_return_lse = do_prefix_merge
|
||||
|
||||
@@ -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.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_context import get_token_to_kv_pool
|
||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.server_args import get_global_server_args
|
||||
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):
|
||||
"""Enable fused set_kv_buffer only on CUDA with bfloat16 KV cache."""
|
||||
pool = get_token_to_kv_pool()
|
||||
return (
|
||||
_is_cuda
|
||||
and hasattr(forward_batch.token_to_kv_pool, "dtype")
|
||||
and forward_batch.token_to_kv_pool.dtype == torch.bfloat16
|
||||
and not isinstance(forward_batch.token_to_kv_pool, SWAKVPool)
|
||||
and pool.dtype == torch.bfloat16
|
||||
and not isinstance(pool, SWAKVPool)
|
||||
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
|
||||
|
||||
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)
|
||||
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_cpu=seq_lens_cpu,
|
||||
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,
|
||||
spec_algorithm=SpeculativeAlgorithm.DFLASH,
|
||||
spec_info=draft_spec_info,
|
||||
|
||||
@@ -24,6 +24,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardBatch,
|
||||
ForwardMode,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
||||
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
||||
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
||||
from sglang.srt.speculative.spec_utils import (
|
||||
@@ -332,8 +333,6 @@ class EAGLEDraftCudaGraphRunner:
|
||||
seq_lens_cpu=seq_lens_cpu,
|
||||
extend_seq_lens=extend_seq_lens,
|
||||
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,
|
||||
seq_lens_sum=seq_lens.sum().item(),
|
||||
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():
|
||||
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,
|
||||
@@ -367,7 +361,6 @@ class EAGLEDraftCudaGraphRunner:
|
||||
)
|
||||
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
|
||||
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)
|
||||
return ret
|
||||
|
||||
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
|
||||
)
|
||||
with forward_context(ForwardContext(attn_backend=self.draft_attn_backend)):
|
||||
self.draft_attn_backend.init_forward_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
|
||||
)
|
||||
|
||||
set_global_graph_memory_pool(graph.pool())
|
||||
return graph, out
|
||||
|
||||
@@ -25,6 +25,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardBatch,
|
||||
ForwardMode,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
||||
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
||||
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
||||
from sglang.srt.speculative.spec_utils import fast_topk
|
||||
@@ -352,8 +353,6 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
num_accept_tokens=num_accept_tokens,
|
||||
)
|
||||
|
||||
self.deepep_adapter.capture(is_extend_in_batch=True)
|
||||
|
||||
# Forward batch
|
||||
forward_batch = ForwardBatch(
|
||||
forward_mode=self.forward_mode,
|
||||
@@ -365,8 +364,6 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
next_token_logits_buffer=next_token_logits_buffer,
|
||||
extend_seq_lens=extend_seq_lens,
|
||||
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,
|
||||
seq_lens_sum=seq_lens.sum().item(),
|
||||
return_logprob=False,
|
||||
@@ -379,23 +376,10 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
spec_algorithm=self.model_runner.spec_algorithm,
|
||||
spec_info=spec_info,
|
||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||
attn_backend=self.draft_extend_attn_backend,
|
||||
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():
|
||||
# model.forward() bypasses _forward_raw(), so invalidate manually.
|
||||
if self.model_runner.is_hybrid_swa:
|
||||
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
|
||||
return ret
|
||||
|
||||
self._capture_init(run_once)
|
||||
|
||||
out = self._capture_graph(
|
||||
graph, get_global_graph_memory_pool(), stream, run_once
|
||||
)
|
||||
with forward_context(
|
||||
ForwardContext(attn_backend=self.draft_extend_attn_backend)
|
||||
):
|
||||
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())
|
||||
return graph, out
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import contextlib
|
||||
import logging
|
||||
import time
|
||||
from contextlib import contextmanager
|
||||
@@ -30,6 +31,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardBatch,
|
||||
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.trace import get_global_tracing_enabled
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
@@ -881,13 +883,18 @@ class EAGLEWorker(TpModelWorker):
|
||||
):
|
||||
out_cache_loc = out_cache_loc.contiguous()
|
||||
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
|
||||
|
||||
# Run forward
|
||||
logits_output = self.draft_model_runner.forward(
|
||||
forward_batch, skip_attn_backend_init=True
|
||||
).logits_output
|
||||
# Run forward under a per-step ForwardContext so the model layer
|
||||
# reads attn_backends[i] for the i-th draft step. ``_forward_raw``
|
||||
# is no-op for the attn_backend half when a context is already
|
||||
# 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}")
|
||||
probs = torch.softmax(logits_output.next_token_logits, 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
|
||||
else:
|
||||
forward_batch.can_run_dp_cuda_graph = False
|
||||
attn_backend = None
|
||||
if not forward_batch.forward_mode.is_idle():
|
||||
attn_backend = (
|
||||
self.draft_extend_attn_backend
|
||||
or self.draft_model_runner.attn_backend
|
||||
)
|
||||
attn_backend.init_forward_metadata(forward_batch)
|
||||
forward_batch.attn_backend = attn_backend
|
||||
logits_output = self.draft_model_runner.forward(
|
||||
forward_batch, skip_attn_backend_init=True
|
||||
).logits_output
|
||||
# Publish the chosen backend via ForwardContext so model code
|
||||
# picks it up for this forward (no runner-attr mutation).
|
||||
if attn_backend is not None:
|
||||
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.
|
||||
probs = torch.softmax(logits_output.next_token_logits, 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.model_executor.cuda_graph_runner import CudaGraphRunner
|
||||
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.speculative.adaptive_runtime_state import (
|
||||
AdaptiveController,
|
||||
@@ -466,13 +467,17 @@ class EagleDraftWorker(BaseDraftWorker):
|
||||
# Set inputs
|
||||
forward_batch.input_ids = input_ids
|
||||
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
|
||||
|
||||
# Run forward
|
||||
logits_output = self.draft_runner.forward(
|
||||
forward_batch, skip_attn_backend_init=True
|
||||
).logits_output
|
||||
# Run forward under a per-step ForwardContext so the model layer
|
||||
# reads attn_backends[i] for the i-th draft step. ``_forward_raw``
|
||||
# honors the outer context and does not override.
|
||||
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}")
|
||||
probs = torch.softmax(logits_output.next_token_logits, 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,
|
||||
ForwardMode,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
||||
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
||||
from sglang.srt.speculative.frozen_kv_mtp_info import FrozenKVMTPDraftInput
|
||||
from sglang.srt.utils import (
|
||||
@@ -266,9 +267,6 @@ class FrozenKVMTPCudaGraphRunner:
|
||||
req_pool_indices=req_pool_indices,
|
||||
seq_lens=seq_lens,
|
||||
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,
|
||||
seq_lens_sum=seq_lens.sum().item(),
|
||||
return_logprob=False,
|
||||
@@ -283,10 +281,6 @@ class FrozenKVMTPCudaGraphRunner:
|
||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||
)
|
||||
|
||||
self.frozen_kv_mtp_worker._init_frozen_kv_metadata_capture_cuda_graph(
|
||||
forward_batch
|
||||
)
|
||||
|
||||
def run_once():
|
||||
if self.model_runner.is_hybrid_swa:
|
||||
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
|
||||
return ret
|
||||
|
||||
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
|
||||
)
|
||||
# Swap the draft backend's token_to_kv_pool to the frozen target pool
|
||||
# for the capture; the single backend-attr swap is seen by both
|
||||
# ``get_token_to_kv_pool()`` (via ``get_attn_backend()``) and the
|
||||
# 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())
|
||||
return graph, out
|
||||
|
||||
|
||||
@@ -14,7 +14,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import contextmanager
|
||||
from typing import Tuple
|
||||
from typing import TYPE_CHECKING, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
@@ -28,39 +28,66 @@ from sglang.srt.speculative.frozen_kv_mtp_info import (
|
||||
)
|
||||
from sglang.srt.speculative.spec_utils import fast_topk
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
|
||||
|
||||
@contextmanager
|
||||
def frozen_kv_target_view(forward_batch: ForwardBatch, kv_context: FrozenKVMTPContext):
|
||||
"""Build attention metadata against committed target-prefix geometry."""
|
||||
def frozen_kv_target_view(
|
||||
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:
|
||||
raise RuntimeError(
|
||||
"Frozen-KV MTP target view called before the model was bound; "
|
||||
"bind the frozen KV context first."
|
||||
)
|
||||
saved_spec_info = forward_batch.spec_info
|
||||
saved_kv_pool = forward_batch.token_to_kv_pool
|
||||
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:
|
||||
yield
|
||||
finally:
|
||||
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
|
||||
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:
|
||||
raise RuntimeError(
|
||||
"Frozen-KV MTP target KV pool view called before the model was bound; "
|
||||
"bind the frozen KV context first."
|
||||
)
|
||||
saved_kv_pool = forward_batch.token_to_kv_pool
|
||||
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:
|
||||
yield
|
||||
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:
|
||||
|
||||
@@ -39,6 +39,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardBatch,
|
||||
ForwardMode,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
||||
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
|
||||
from sglang.srt.observability.req_time_stats import set_time_batch
|
||||
from sglang.srt.observability.trace import get_global_tracing_enabled
|
||||
@@ -248,10 +249,14 @@ class FrozenKVMTPWorker(TpModelWorker):
|
||||
self.kv_context = ctx
|
||||
|
||||
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):
|
||||
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:
|
||||
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()
|
||||
with self._frozen_kv_target_view(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(
|
||||
self, forward_batch: ForwardBatch
|
||||
@@ -290,7 +294,6 @@ class FrozenKVMTPWorker(TpModelWorker):
|
||||
forward_mode=ForwardMode.DECODE,
|
||||
spec_info=None,
|
||||
)
|
||||
forward_batch.attn_backend = self.draft_attn_backend
|
||||
|
||||
def _init_frozen_kv_metadata_replay_cuda_graph(
|
||||
self, forward_batch: ForwardBatch, bs: int, seq_lens_sum: int
|
||||
@@ -310,7 +313,6 @@ class FrozenKVMTPWorker(TpModelWorker):
|
||||
else None
|
||||
),
|
||||
)
|
||||
forward_batch.attn_backend = self.draft_attn_backend
|
||||
|
||||
def init_cuda_graphs(self) -> None:
|
||||
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
|
||||
self._set_positions(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(
|
||||
forward_batch, skip_attn_backend_init=True
|
||||
).logits_output
|
||||
@@ -678,7 +682,9 @@ class FrozenKVMTPWorker(TpModelWorker):
|
||||
forward_batch.spec_info.hidden_states = hidden_states
|
||||
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(
|
||||
forward_batch, skip_attn_backend_init=True
|
||||
).logits_output
|
||||
|
||||
@@ -40,6 +40,11 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardBatch,
|
||||
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.speculative.eagle_info import EagleDraftExtendInput
|
||||
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_cpu=seq_lens_cpu,
|
||||
next_token_logits_buffer=next_token_logits_buffer,
|
||||
req_to_token_pool=self.model_runner.req_to_token_pool,
|
||||
token_to_kv_pool=self.model_runner.token_to_kv_pool,
|
||||
out_cache_loc=out_cache_loc,
|
||||
seq_lens_sum=seq_lens.sum().item(),
|
||||
return_logprob=False,
|
||||
@@ -383,7 +386,6 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
spec_algorithm=self.model_runner.spec_algorithm,
|
||||
spec_info=spec_info,
|
||||
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_cpu=extend_seq_lens_cpu,
|
||||
padded_static_len=self.padded_static_len,
|
||||
@@ -400,26 +402,11 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
graph = self._create_graph()
|
||||
stream = self.stream
|
||||
|
||||
self.deepep_adapter.capture(is_extend_in_batch=True)
|
||||
|
||||
num_tokens = bs * self.num_tokens_per_bs
|
||||
forward_batch = self.get_forward_batch(bs)
|
||||
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():
|
||||
# model.forward() bypasses _forward_raw(), so invalidate manually.
|
||||
if self.model_runner.is_hybrid_swa:
|
||||
self.model_runner.token_to_kv_pool.invalidate_loc_cache()
|
||||
|
||||
@@ -490,18 +477,28 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
forward_batch.batch_size,
|
||||
self.step,
|
||||
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,
|
||||
)
|
||||
forward_batch.out_cache_loc = output_cache_loc_backup
|
||||
forward_batch.spec_info.hidden_states = hidden_states_backup
|
||||
return ret
|
||||
|
||||
self._capture_init(run_once)
|
||||
|
||||
out = self._capture_graph(
|
||||
graph, get_global_graph_memory_pool(), stream, run_once
|
||||
)
|
||||
with forward_context(ForwardContext(attn_backend=attn_backend)):
|
||||
attn_backend.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,
|
||||
)
|
||||
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())
|
||||
return graph, out
|
||||
|
||||
@@ -419,9 +419,6 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
|
||||
topk_p_list = []
|
||||
topk_index_list = []
|
||||
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(
|
||||
forward_batch
|
||||
)
|
||||
@@ -526,9 +523,6 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
|
||||
draft_logits_output.topk_index,
|
||||
)
|
||||
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(
|
||||
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.mem_cache.memory_pool import MHATokenToKVPool
|
||||
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
|
||||
|
||||
|
||||
@@ -109,6 +113,9 @@ class TestFlashAttentionBackend(CustomTestCase):
|
||||
self.backend = FlashAttentionBackend(self.model_runner)
|
||||
self.ref_backend = TorchNativeAttnBackend(self.model_runner)
|
||||
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):
|
||||
# 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(
|
||||
[q_len] * self.batch_size, device="cpu"
|
||||
),
|
||||
attn_backend=self.backend,
|
||||
)
|
||||
if attn_cp_size > 1:
|
||||
forward_batch.attn_cp_metadata = type(
|
||||
@@ -273,16 +279,11 @@ class TestFlashAttentionBackend(CustomTestCase):
|
||||
[total_len] * self.batch_size, device=self.device
|
||||
),
|
||||
seq_lens_cpu=torch.tensor([total_len] * self.batch_size, device="cpu"),
|
||||
attn_backend=self.backend,
|
||||
)
|
||||
|
||||
# Add token pool
|
||||
forward_batch.req_to_token_pool = self.model_runner.req_to_token_pool
|
||||
|
||||
# Write current batch's req_to_token to req_to_token_pool
|
||||
# Pool refs are resolved via the active ForwardContext (published in
|
||||
# setUp). Write the test fixture's req_to_token mapping.
|
||||
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
|
||||
|
||||
@@ -307,7 +308,7 @@ class TestFlashAttentionBackend(CustomTestCase):
|
||||
)
|
||||
|
||||
# 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,
|
||||
torch.arange(self.batch_size * cache_len, device=self.device),
|
||||
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.mem_cache.memory_pool import MLATokenToKVPool
|
||||
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
|
||||
|
||||
|
||||
@@ -112,6 +116,8 @@ class TestFlashAttentionMLABackend(CustomTestCase):
|
||||
self.backend = FlashAttentionBackend(self.model_runner)
|
||||
self.ref_backend = TorchNativeAttnBackend(self.model_runner)
|
||||
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):
|
||||
self.model_runner = MockModelRunner(
|
||||
@@ -192,7 +198,6 @@ class TestFlashAttentionMLABackend(CustomTestCase):
|
||||
extend_seq_lens_cpu=torch.tensor(
|
||||
[q_len] * self.batch_size, device="cpu"
|
||||
),
|
||||
attn_backend=self.backend,
|
||||
)
|
||||
|
||||
else: # ForwardMode.DECODE
|
||||
@@ -216,15 +221,10 @@ class TestFlashAttentionMLABackend(CustomTestCase):
|
||||
[total_len] * self.batch_size, device=self.device
|
||||
),
|
||||
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
|
||||
forward_batch.req_to_token_pool = self.model_runner.req_to_token_pool
|
||||
|
||||
# Add KV cache from model runner to forward batch
|
||||
forward_batch.token_to_kv_pool = self.model_runner.token_to_kv_pool
|
||||
|
||||
# Pool refs are resolved via the active ForwardContext (published in
|
||||
# setUp); the fixture no longer needs to attach them to forward_batch.
|
||||
return forward_batch
|
||||
|
||||
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
|
||||
forward_batch.token_to_kv_pool.set_mla_kv_buffer(
|
||||
self.model_runner.token_to_kv_pool.set_mla_kv_buffer(
|
||||
layer,
|
||||
torch.arange(self.batch_size * cache_len, device=self.device),
|
||||
cache_k_nope,
|
||||
|
||||
@@ -110,12 +110,10 @@ class MockReqToTokenPool:
|
||||
|
||||
|
||||
# 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):
|
||||
computed_kv_indices = forward_batch.prefix_chunk_kv_indices[i]
|
||||
req_to_token = forward_batch.req_to_token_pool.req_to_token[
|
||||
: forward_batch.batch_size, :
|
||||
]
|
||||
req_to_token = req_to_token_pool.req_to_token[: forward_batch.batch_size, :]
|
||||
ref_kv_indices = torch.empty(
|
||||
forward_batch.prefix_chunk_num_tokens[i],
|
||||
dtype=torch.int32,
|
||||
@@ -205,8 +203,20 @@ class TestPrefixChunkInfo(CustomTestCase):
|
||||
extend_prefix_lens=prefix_lens,
|
||||
extend_prefix_lens_cpu=prefix_lens_cpu,
|
||||
)
|
||||
forward_batch.req_to_token_pool = self.req_to_token_pool
|
||||
forward_batch.token_to_kv_pool = self.token_to_kv_pool
|
||||
# Pool refs are resolved via the active ForwardContext; mock an
|
||||
# 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)
|
||||
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),
|
||||
)
|
||||
|
||||
check_kv_indices(forward_batch)
|
||||
check_kv_indices(forward_batch, self.req_to_token_pool)
|
||||
|
||||
|
||||
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.mem_cache.memory_pool import MLATokenToKVPool
|
||||
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 (
|
||||
ServerArgs,
|
||||
get_global_server_args,
|
||||
@@ -434,10 +438,9 @@ class TestTRTLLMMLA(CustomTestCase):
|
||||
req_pool_indices=torch.arange(batch_size, device=config["device"]),
|
||||
seq_lens=seq_lens,
|
||||
seq_lens_cpu=seq_lens.cpu(),
|
||||
attn_backend=backend,
|
||||
)
|
||||
fb.req_to_token_pool = model_runner.req_to_token_pool
|
||||
fb.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||
# Publish backend for RadixAttention dispatch.
|
||||
set_forward_context(ForwardContext(attn_backend=backend))
|
||||
|
||||
# Add position information for RoPE
|
||||
fb.positions = torch.arange(batch_size, device=config["device"])
|
||||
@@ -1167,10 +1170,9 @@ class TestTRTLLMMLA(CustomTestCase):
|
||||
seq_lens_cpu=seq_lens.cpu(),
|
||||
attn_attend_prefix_cache=False,
|
||||
mha_return_lse=False,
|
||||
attn_backend=backend,
|
||||
)
|
||||
fb.req_to_token_pool = model_runner.req_to_token_pool
|
||||
fb.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||
# Publish backend for RadixAttention dispatch.
|
||||
set_forward_context(ForwardContext(attn_backend=backend))
|
||||
|
||||
# Add position information for RoPE
|
||||
fb.positions = torch.arange(batch_size, device=config["device"])
|
||||
|
||||
@@ -360,7 +360,6 @@ class TestDSAIndexer(CustomTestCase):
|
||||
),
|
||||
extend_seq_lens=torch.tensor([q_len] * batch_size, device=self.device),
|
||||
extend_seq_lens_cpu=torch.tensor([q_len] * batch_size, device="cpu"),
|
||||
attn_backend=self.backend,
|
||||
)
|
||||
else: # ForwardMode.DECODE
|
||||
decode_len = 1
|
||||
@@ -379,12 +378,18 @@ class TestDSAIndexer(CustomTestCase):
|
||||
req_pool_indices=torch.arange(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"),
|
||||
attn_backend=self.backend,
|
||||
)
|
||||
|
||||
# Add token pools
|
||||
forward_batch.req_to_token_pool = self.model_runner.req_to_token_pool
|
||||
forward_batch.token_to_kv_pool = self.model_runner.token_to_kv_pool
|
||||
# Pool refs + attn_backend are now resolved via the ForwardContext;
|
||||
# publish ``self.backend`` for the duration of this fixture call so
|
||||
# ``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
|
||||
page_size = self.model_runner.page_size
|
||||
|
||||
Reference in New Issue
Block a user