From c5251a98a9d499d600beb557835ac5874e0c3f36 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Thu, 21 May 2026 14:01:49 -0700 Subject: [PATCH] feat(model_runner): remove pool/backend refs from ForwardBatch via ForwardContext (#25983) Co-authored-by: Claude Sonnet 4.6 (1M context) --- python/sglang/srt/batch_overlap/operations.py | 69 +++++- .../srt/batch_overlap/two_batch_overlap.py | 14 +- .../musa/attention/flashattention_backend.py | 26 +-- .../npu/attention/ascend_backend.py | 123 ++++++----- .../npu/attention/mla_preprocess.py | 14 +- .../modules/deepseek_v2_attention_mla_npu.py | 7 +- .../srt/layers/attention/aiter_backend.py | 47 ++-- .../layers/attention/cutlass_mla_backend.py | 6 +- .../layers/attention/deepseek_v4_backend.py | 19 +- .../deepseek_v4_backend_hip_radix.py | 19 +- .../srt/layers/attention/dsa/dsa_indexer.py | 87 ++++---- .../srt/layers/attention/dsa_backend.py | 55 +++-- .../srt/layers/attention/dsv4/compress_hip.py | 57 +++-- .../srt/layers/attention/dsv4/compressor.py | 32 +-- .../layers/attention/dsv4/compressor_v2.py | 4 +- .../srt/layers/attention/dsv4/indexer.py | 8 +- .../dual_chunk_flashattention_backend.py | 18 +- .../attention/flashattention_backend.py | 50 ++--- .../layers/attention/flashinfer_backend.py | 24 +- .../attention/flashinfer_mla_backend.py | 25 +-- .../srt/layers/attention/flashmla_backend.py | 8 +- .../layers/attention/hybrid_attn_backend.py | 2 + .../attention/hybrid_linear_attn_backend.py | 4 + .../srt/layers/attention/intel_amx_backend.py | 18 +- .../srt/layers/attention/tbo_backend.py | 4 + .../attention/tokenspeed_mla_backend.py | 2 +- .../layers/attention/torch_flex_backend.py | 20 +- .../layers/attention/torch_native_backend.py | 23 +- .../srt/layers/attention/triton_backend.py | 33 +-- .../layers/attention/trtllm_mha_backend.py | 18 +- .../layers/attention/trtllm_mla_backend.py | 8 +- .../srt/layers/attention/wave_backend.py | 17 +- .../srt/layers/attention/xpu_backend.py | 48 ++-- python/sglang/srt/layers/radix_attention.py | 5 +- .../srt/layers/radix_linear_attention.py | 5 +- python/sglang/srt/layers/utils/cp_utils.py | 3 +- .../breakable_cuda_graph_runner.py | 28 ++- .../srt/model_executor/cpu_graph_runner.py | 77 +++---- .../srt/model_executor/cuda_graph_runner.py | 143 ++++++------ .../forward_batch_deepseek_mha_mixin.py | 16 +- .../srt/model_executor/forward_batch_info.py | 14 -- .../srt/model_executor/forward_context.py | 84 +++++++ .../sglang/srt/model_executor/model_runner.py | 205 ++++++++++-------- .../piecewise_cuda_graph_runner.py | 118 +++++----- .../attention_backend_handler.py | 3 +- .../attention_forward_methods/forward_mha.py | 32 +-- .../attention_forward_methods/forward_mla.py | 16 +- .../forward_mla_fused_rope_rocm.py | 16 +- python/sglang/srt/models/deepseek_v4.py | 25 ++- python/sglang/srt/models/deepseek_v4_nextn.py | 8 +- python/sglang/srt/models/falcon_h1.py | 3 +- python/sglang/srt/models/gemma3_mm.py | 7 +- python/sglang/srt/models/gemma4_mm.py | 9 +- python/sglang/srt/models/granitemoehybrid.py | 3 +- python/sglang/srt/models/jet_nemotron.py | 9 +- python/sglang/srt/models/lfm2.py | 7 +- python/sglang/srt/models/lfm2_moe.py | 7 +- python/sglang/srt/models/mindspore.py | 12 +- python/sglang/srt/models/nemotron_h.py | 5 +- python/sglang/srt/models/qwen3.py | 3 +- python/sglang/srt/models/sarvam_moe.py | 12 +- python/sglang/srt/models/utils.py | 9 +- .../sglang/srt/speculative/dflash_worker.py | 3 - .../eagle_draft_cuda_graph_runner.py | 25 +-- .../eagle_draft_extend_cuda_graph_runner.py | 40 ++-- python/sglang/srt/speculative/eagle_worker.py | 32 ++- .../sglang/srt/speculative/eagle_worker_v2.py | 15 +- .../frozen_kv_mtp_cuda_graph_runner.py | 32 ++- .../srt/speculative/frozen_kv_mtp_utils.py | 47 +++- .../srt/speculative/frozen_kv_mtp_worker.py | 20 +- ...er_eagle_draft_extend_cuda_graph_runner.py | 47 ++-- .../multi_layer_eagle_worker_v2.py | 6 - .../attention/test_flashattn_backend.py | 19 +- .../attention/test_flashattn_mla_backend.py | 18 +- .../attention/test_prefix_chunk_info.py | 24 +- .../attention/test_trtllm_mla_backend.py | 14 +- test/registered/kernels/test_dsa_indexer.py | 15 +- 77 files changed, 1236 insertions(+), 914 deletions(-) create mode 100644 python/sglang/srt/model_executor/forward_context.py diff --git a/python/sglang/srt/batch_overlap/operations.py b/python/sglang/srt/batch_overlap/operations.py index 3d61ac82f..729e00bcb 100644 --- a/python/sglang/srt/batch_overlap/operations.py +++ b/python/sglang/srt/batch_overlap/operations.py @@ -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( diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index f351851d5..de9faca4e 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -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, diff --git a/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py b/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py index 17fb35ae6..6044e0a82 100644 --- a/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py +++ b/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py @@ -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( diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py index 1035a96eb..038119688 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py @@ -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) diff --git a/python/sglang/srt/hardware_backend/npu/attention/mla_preprocess.py b/python/sglang/srt/hardware_backend/npu/attention/mla_preprocess.py index 51cf7421e..1107f11b2 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/mla_preprocess.py +++ b/python/sglang/srt/hardware_backend/npu/attention/mla_preprocess.py @@ -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, diff --git a/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py b/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py index 68f23f1ac..79f0bb86a 100644 --- a/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py +++ b/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py @@ -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 ) diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 897740536..57f45c2a5 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -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, diff --git a/python/sglang/srt/layers/attention/cutlass_mla_backend.py b/python/sglang/srt/layers/attention/cutlass_mla_backend.py index e81e761bc..05641ea1f 100644 --- a/python/sglang/srt/layers/attention/cutlass_mla_backend.py +++ b/python/sglang/srt/layers/attention/cutlass_mla_backend.py @@ -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, diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index f9f396428..bf743a2c2 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -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): diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py index 9a9a7225a..7b8c9dc59 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py @@ -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): diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index 6ffb9719b..7548c811f 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -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 diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index a2c062b91..9e5c8f134 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -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( diff --git a/python/sglang/srt/layers/attention/dsv4/compress_hip.py b/python/sglang/srt/layers/attention/dsv4/compress_hip.py index 1c69f7e46..fa70fc693 100644 --- a/python/sglang/srt/layers/attention/dsv4/compress_hip.py +++ b/python/sglang/srt/layers/attention/dsv4/compress_hip.py @@ -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 + ) diff --git a/python/sglang/srt/layers/attention/dsv4/compressor.py b/python/sglang/srt/layers/attention/dsv4/compressor.py index 092b98e2c..e663c6ab8 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor.py @@ -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), diff --git a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py index 5d6dd1e0d..13b5c65cc 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py @@ -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) diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index f8899264d..5284a240f 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -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, diff --git a/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py b/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py index a84015a80..fa0da5c46 100644 --- a/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py +++ b/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py @@ -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, diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index cc3b1ca32..aa8eb009e 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -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( diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index 27705a4b8..13930a752 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -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, diff --git a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py index 601a80cea..61b6c49a5 100644 --- a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py @@ -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, diff --git a/python/sglang/srt/layers/attention/flashmla_backend.py b/python/sglang/srt/layers/attention/flashmla_backend.py index 1e63f9b5c..c0bce60ce 100644 --- a/python/sglang/srt/layers/attention/flashmla_backend.py +++ b/python/sglang/srt/layers/attention/flashmla_backend.py @@ -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: diff --git a/python/sglang/srt/layers/attention/hybrid_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_attn_backend.py index 69e80149e..df0c70dc5 100644 --- a/python/sglang/srt/layers/attention/hybrid_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_attn_backend.py @@ -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: """ diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index 7b41b66e6..a388e9d33 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -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 diff --git a/python/sglang/srt/layers/attention/intel_amx_backend.py b/python/sglang/srt/layers/attention/intel_amx_backend.py index 46b657d64..2f5e4141b 100644 --- a/python/sglang/srt/layers/attention/intel_amx_backend.py +++ b/python/sglang/srt/layers/attention/intel_amx_backend.py @@ -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, diff --git a/python/sglang/srt/layers/attention/tbo_backend.py b/python/sglang/srt/layers/attention/tbo_backend.py index 2ae120686..76d83b7b7 100644 --- a/python/sglang/srt/layers/attention/tbo_backend.py +++ b/python/sglang/srt/layers/attention/tbo_backend.py @@ -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]): diff --git a/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py index 9ab17576a..5e513f081 100644 --- a/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py +++ b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py @@ -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), diff --git a/python/sglang/srt/layers/attention/torch_flex_backend.py b/python/sglang/srt/layers/attention/torch_flex_backend.py index 69f097efd..1af8508cb 100644 --- a/python/sglang/srt/layers/attention/torch_flex_backend.py +++ b/python/sglang/srt/layers/attention/torch_flex_backend.py @@ -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, diff --git a/python/sglang/srt/layers/attention/torch_native_backend.py b/python/sglang/srt/layers/attention/torch_native_backend.py index 00d424c44..8894f92ca 100644 --- a/python/sglang/srt/layers/attention/torch_native_backend.py +++ b/python/sglang/srt/layers/attention/torch_native_backend.py @@ -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, diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index 206037f49..ea727713f 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -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, diff --git a/python/sglang/srt/layers/attention/trtllm_mha_backend.py b/python/sglang/srt/layers/attention/trtllm_mha_backend.py index 722c71022..e68bcb95e 100644 --- a/python/sglang/srt/layers/attention/trtllm_mha_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mha_backend.py @@ -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) diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index 87f9c281d..72b8e6d5f 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -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) diff --git a/python/sglang/srt/layers/attention/wave_backend.py b/python/sglang/srt/layers/attention/wave_backend.py index 2a759c222..ebdadc9bf 100644 --- a/python/sglang/srt/layers/attention/wave_backend.py +++ b/python/sglang/srt/layers/attention/wave_backend.py @@ -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, diff --git a/python/sglang/srt/layers/attention/xpu_backend.py b/python/sglang/srt/layers/attention/xpu_backend.py index e918af462..c7149c1a1 100644 --- a/python/sglang/srt/layers/attention/xpu_backend.py +++ b/python/sglang/srt/layers/attention/xpu_backend.py @@ -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: diff --git a/python/sglang/srt/layers/radix_attention.py b/python/sglang/srt/layers/radix_attention.py index 468766011..1e8784f1d 100644 --- a/python/sglang/srt/layers/radix_attention.py +++ b/python/sglang/srt/layers/radix_attention.py @@ -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, diff --git a/python/sglang/srt/layers/radix_linear_attention.py b/python/sglang/srt/layers/radix_linear_attention.py index edaac2125..ad07aebe9 100644 --- a/python/sglang/srt/layers/radix_linear_attention.py +++ b/python/sglang/srt/layers/radix_linear_attention.py @@ -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], diff --git a/python/sglang/srt/layers/utils/cp_utils.py b/python/sglang/srt/layers/utils/cp_utils.py index 885dfed3b..e998dfc67 100644 --- a/python/sglang/srt/layers/utils/cp_utils.py +++ b/python/sglang/srt/layers/utils/cp_utils.py @@ -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, diff --git a/python/sglang/srt/model_executor/breakable_cuda_graph_runner.py b/python/sglang/srt/model_executor/breakable_cuda_graph_runner.py index 1365d80fe..7d2a495c9 100644 --- a/python/sglang/srt/model_executor/breakable_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/breakable_cuda_graph_runner.py @@ -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 diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index edf0dadcb..b80043e70 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -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): diff --git a/python/sglang/srt/model_executor/cuda_graph_runner.py b/python/sglang/srt/model_executor/cuda_graph_runner.py index c2e6d121f..65e38c2d0 100644 --- a/python/sglang/srt/model_executor/cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/cuda_graph_runner.py @@ -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 diff --git a/python/sglang/srt/model_executor/forward_batch_deepseek_mha_mixin.py b/python/sglang/srt/model_executor/forward_batch_deepseek_mha_mixin.py index 769d30235..14b427b33 100644 --- a/python/sglang/srt/model_executor/forward_batch_deepseek_mha_mixin.py +++ b/python/sglang/srt/model_executor/forward_batch_deepseek_mha_mixin.py @@ -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 diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 89f88c235..6698429aa 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -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, diff --git a/python/sglang/srt/model_executor/forward_context.py b/python/sglang/srt/model_executor/forward_context.py new file mode 100644 index 000000000..3a3a7e50f --- /dev/null +++ b/python/sglang/srt/model_executor/forward_context.py @@ -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) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 55850e5ea..52a483802 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -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 ): diff --git a/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py b/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py index 39a516a45..8653bce86 100644 --- a/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py @@ -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, diff --git a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py index de8c6b322..e4cac55f9 100644 --- a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py +++ b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py @@ -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: diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py index 75019ba11..67f4483cd 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py @@ -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) diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index e0ad07511..904f0ee03 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -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: diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_fused_rope_rocm.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_fused_rope_rocm.py index 8868897af..65545069e 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_fused_rope_rocm.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_fused_rope_rocm.py @@ -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 ) diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index fdf6d557c..4b062a119 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -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): diff --git a/python/sglang/srt/models/deepseek_v4_nextn.py b/python/sglang/srt/models/deepseek_v4_nextn.py index f6d9a3f7d..1260320a8 100644 --- a/python/sglang/srt/models/deepseek_v4_nextn.py +++ b/python/sglang/srt/models/deepseek_v4_nextn.py @@ -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) diff --git a/python/sglang/srt/models/falcon_h1.py b/python/sglang/srt/models/falcon_h1.py index 72f684c2b..b12986658 100644 --- a/python/sglang/srt/models/falcon_h1.py +++ b/python/sglang/srt/models/falcon_h1.py @@ -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 diff --git a/python/sglang/srt/models/gemma3_mm.py b/python/sglang/srt/models/gemma3_mm.py index 25745d331..9b362dbba 100644 --- a/python/sglang/srt/models/gemma3_mm.py +++ b/python/sglang/srt/models/gemma3_mm.py @@ -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 ) diff --git a/python/sglang/srt/models/gemma4_mm.py b/python/sglang/srt/models/gemma4_mm.py index fb14dd17a..cafc31f20 100644 --- a/python/sglang/srt/models/gemma4_mm.py +++ b/python/sglang/srt/models/gemma4_mm.py @@ -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 diff --git a/python/sglang/srt/models/granitemoehybrid.py b/python/sglang/srt/models/granitemoehybrid.py index e18aeb466..a26e71601 100644 --- a/python/sglang/srt/models/granitemoehybrid.py +++ b/python/sglang/srt/models/granitemoehybrid.py @@ -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( diff --git a/python/sglang/srt/models/jet_nemotron.py b/python/sglang/srt/models/jet_nemotron.py index 1e6d2ec87..fec8c1fb6 100644 --- a/python/sglang/srt/models/jet_nemotron.py +++ b/python/sglang/srt/models/jet_nemotron.py @@ -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 diff --git a/python/sglang/srt/models/lfm2.py b/python/sglang/srt/models/lfm2.py index 2750a0f81..1f4f7544e 100644 --- a/python/sglang/srt/models/lfm2.py +++ b/python/sglang/srt/models/lfm2.py @@ -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) diff --git a/python/sglang/srt/models/lfm2_moe.py b/python/sglang/srt/models/lfm2_moe.py index fcc396357..4846b0b99 100644 --- a/python/sglang/srt/models/lfm2_moe.py +++ b/python/sglang/srt/models/lfm2_moe.py @@ -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) diff --git a/python/sglang/srt/models/mindspore.py b/python/sglang/srt/models/mindspore.py index da95ab139..b91197286 100644 --- a/python/sglang/srt/models/mindspore.py +++ b/python/sglang/srt/models/mindspore.py @@ -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 diff --git a/python/sglang/srt/models/nemotron_h.py b/python/sglang/srt/models/nemotron_h.py index 1e879455f..b60ffa372 100644 --- a/python/sglang/srt/models/nemotron_h.py +++ b/python/sglang/srt/models/nemotron_h.py @@ -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 diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index 21d262b71..2a18e2adb 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -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 diff --git a/python/sglang/srt/models/sarvam_moe.py b/python/sglang/srt/models/sarvam_moe.py index bca26936d..6ef737727 100644 --- a/python/sglang/srt/models/sarvam_moe.py +++ b/python/sglang/srt/models/sarvam_moe.py @@ -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 diff --git a/python/sglang/srt/models/utils.py b/python/sglang/srt/models/utils.py index 92588e177..341f7b458 100644 --- a/python/sglang/srt/models/utils.py +++ b/python/sglang/srt/models/utils.py @@ -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) diff --git a/python/sglang/srt/speculative/dflash_worker.py b/python/sglang/srt/speculative/dflash_worker.py index 87ddcfe23..86cd76bf7 100644 --- a/python/sglang/srt/speculative/dflash_worker.py +++ b/python/sglang/srt/speculative/dflash_worker.py @@ -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, diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index dc379878b..c0d180231 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -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 diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index 23e79648f..9df3742fc 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -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 diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index 88488bf88..3f3964d57 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -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) diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 4a9969113..3f24eebe4 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -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) diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py index da69b3cbd..8b1ac37f8 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py @@ -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 diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_utils.py b/python/sglang/srt/speculative/frozen_kv_mtp_utils.py index 043d8b63f..dbd63c2e4 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_utils.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_utils.py @@ -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: diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_worker.py b/python/sglang/srt/speculative/frozen_kv_mtp_worker.py index 6e4ecdf03..4bad85187 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_worker.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_worker.py @@ -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 diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index c84f8c1e8..30beb43b3 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -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 diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index c8da2727e..aebd5d0ab 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -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 ) diff --git a/test/manual/attention/test_flashattn_backend.py b/test/manual/attention/test_flashattn_backend.py index 16b7b68fd..863d5e36e 100644 --- a/test/manual/attention/test_flashattn_backend.py +++ b/test/manual/attention/test_flashattn_backend.py @@ -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, diff --git a/test/manual/attention/test_flashattn_mla_backend.py b/test/manual/attention/test_flashattn_mla_backend.py index 98eaa5913..f1d53fcf9 100644 --- a/test/manual/attention/test_flashattn_mla_backend.py +++ b/test/manual/attention/test_flashattn_mla_backend.py @@ -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, diff --git a/test/manual/attention/test_prefix_chunk_info.py b/test/manual/attention/test_prefix_chunk_info.py index 5002a0b09..523e6b20f 100644 --- a/test/manual/attention/test_prefix_chunk_info.py +++ b/test/manual/attention/test_prefix_chunk_info.py @@ -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__": diff --git a/test/manual/attention/test_trtllm_mla_backend.py b/test/manual/attention/test_trtllm_mla_backend.py index 6ba9a14c0..e800ed582 100755 --- a/test/manual/attention/test_trtllm_mla_backend.py +++ b/test/manual/attention/test_trtllm_mla_backend.py @@ -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"]) diff --git a/test/registered/kernels/test_dsa_indexer.py b/test/registered/kernels/test_dsa_indexer.py index 09021180f..979d9e60d 100644 --- a/test/registered/kernels/test_dsa_indexer.py +++ b/test/registered/kernels/test_dsa_indexer.py @@ -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