From 6953dae0057b197cad4b665bba6692d2ce02516c Mon Sep 17 00:00:00 2001 From: Qiaolin Yu Date: Sat, 12 Sep 2026 01:32:01 -0700 Subject: [PATCH] [Qwen 3.8 Next] Remove unused tokenwise QSA implementation and tests (#38960) --- .../jit/csrc/elementwise/fast_topk.cuh | 6 +- .../arg_groups/model_overrides/qwen4_exp.py | 3 +- .../srt/layers/attention/qsa/__init__.py | 5 - .../sglang/srt/layers/attention/qsa/config.py | 88 +---- .../srt/layers/attention/qsa/dsa_indexer.py | 337 ------------------ .../sglang/srt/layers/attention/qsa/glue.py | 24 +- .../attention/qwen_sparse_attn_backend.py | 128 +++---- .../srt/mem_cache/kv_cache_configurator.py | 9 - python/sglang/srt/mem_cache/qsa_kv_pool.py | 104 ------ .../srt/model_executor/pool_configurator.py | 12 +- python/sglang/srt/speculative/draft_utils.py | 6 +- test/registered/kernel/qsa/test_qsa.py | 40 +-- 12 files changed, 71 insertions(+), 691 deletions(-) delete mode 100644 python/sglang/srt/layers/attention/qsa/dsa_indexer.py diff --git a/python/sglang/kernels/jit/csrc/elementwise/fast_topk.cuh b/python/sglang/kernels/jit/csrc/elementwise/fast_topk.cuh index 4b015b2f0..b30f4636c 100644 --- a/python/sglang/kernels/jit/csrc/elementwise/fast_topk.cuh +++ b/python/sglang/kernels/jit/csrc/elementwise/fast_topk.cuh @@ -39,7 +39,7 @@ SGL_DEVICE auto convert_to_uint32(float x) -> uint32_t { // When length <= kTopK, write the indices directly. template -SGL_DEVICE void naive_topk(const float* __restrict__ score, int32_t* __restrict__ indice, int32_t length) { +SGL_DEVICE void naive_topk(int32_t* __restrict__ indice, int32_t length) { const auto tid = threadIdx.x; for (int i = tid; i < kTopK; i += kThreadsPerBlock) { indice[i] = (i < length) ? i : -1; @@ -217,12 +217,12 @@ __global__ __launch_bounds__(fast_topk_detail::kThreadsPerBlock) void fast_topk_ device::PDLWaitPrimary(); const auto bid = static_cast(blockIdx.x); - const auto row_start = params.row_starts == nullptr ? 0 : params.row_starts[bid]; + const auto row_start = params.row_starts[bid]; const auto length = params.lengths[bid]; const auto indice = params.indices + bid * kTopK; const auto score = params.input + bid * params.input_stride; if (length <= kTopK) { - naive_topk(score, indice, length); + naive_topk(indice, length); } else { radix_select_topk(score, indice, row_start, length); } diff --git a/python/sglang/srt/arg_groups/model_overrides/qwen4_exp.py b/python/sglang/srt/arg_groups/model_overrides/qwen4_exp.py index 8cac79f4f..2c96fc22f 100644 --- a/python/sglang/srt/arg_groups/model_overrides/qwen4_exp.py +++ b/python/sglang/srt/arg_groups/model_overrides/qwen4_exp.py @@ -74,12 +74,11 @@ def _qwen4_exp_overrides(server_args: Any, hf_config: Any) -> dict: overrides["page_size"] = 64 if sm100_default_attn_backend == "trtllm_mha" else 1 from sglang.srt.layers.attention.qsa.config import ( - QSA_VARIANT_COMPRESSED, parse_qsa_profile, ) profile = parse_qsa_profile(hf_config) - if profile is not None and profile.variant == QSA_VARIANT_COMPRESSED: + if profile is not None: # Compressed slot = full_slot // ratio; all backends need page-aligned pages. # mamba_radix_cache_strategy resolves later, so do not gate on it. overrides["page_size"] = 64 diff --git a/python/sglang/srt/layers/attention/qsa/__init__.py b/python/sglang/srt/layers/attention/qsa/__init__.py index d7a9e3b12..912ec4122 100644 --- a/python/sglang/srt/layers/attention/qsa/__init__.py +++ b/python/sglang/srt/layers/attention/qsa/__init__.py @@ -8,7 +8,6 @@ __all__ = [ "QSAIndexer", "QSAIndexerMetadata", "QSAProfile", - "QwenDSAIndexer", "build_qsa_indexer", "get_qsa_indexer_metadata", "is_qwen_qsa", @@ -21,10 +20,6 @@ def __getattr__(name): from sglang.srt.layers.attention.qsa.qsa_indexer import QSAIndexer return QSAIndexer - if name == "QwenDSAIndexer": - from sglang.srt.layers.attention.qsa.dsa_indexer import QwenDSAIndexer - - return QwenDSAIndexer if name == "QSAIndexerMetadata": from sglang.srt.layers.attention.qsa.metadata import QSAIndexerMetadata diff --git a/python/sglang/srt/layers/attention/qsa/config.py b/python/sglang/srt/layers/attention/qsa/config.py index 4e0c0c954..42df016e1 100644 --- a/python/sglang/srt/layers/attention/qsa/config.py +++ b/python/sglang/srt/layers/attention/qsa/config.py @@ -1,12 +1,4 @@ -"""Shared QSA profile parsing across model variants. - -``QSAProfile`` normalizes each model family's HF-config indexer schema, -so backends, draft utilities and model glue branch on a stable variant name, -not on raw config keys. ``compressed`` is Qwen4-Exp block compression; -``tokenwise`` is qsa_0511 / Qwen3.5-DSA per-token indexing. -DeepSeek NSA configs also expose ``index_topk``, -so the tokenwise schema is additionally gated on a Qwen ``model_type``. -""" +"""QSA profile parsing for Qwen4-Exp compressed indexing.""" from __future__ import annotations @@ -14,14 +6,6 @@ from typing import Optional import msgspec -# QSA variant names. -QSA_VARIANT_COMPRESSED = "compressed" -QSA_VARIANT_TOKENWISE = "tokenwise" - -# Rotary layouts the indexer can consume. -QSA_ROPE_MROPE = "mrope" -QSA_ROPE_PLAIN = "plain" - _COMPRESSED_FIELDS = ( "indexer_n_heads", "indexer_kv_heads", @@ -29,33 +13,22 @@ _COMPRESSED_FIELDS = ( "indexer_budget", "indexer_compress_ratio", ) -_TOKENWISE_FIELDS = ( - "index_topk", - "index_n_heads", - "index_kv_heads", - "index_head_dim", -) - # fast_topk_v2 only supports these compressed block top-k widths. _COMPRESSED_BLOCK_TOPK = frozenset({512, 2048}) -# fast_topk_v2 only supports a 2048-wide tokenwise top-k. -_TOKENWISE_BUDGET = 2048 class QSAProfile(msgspec.Struct, frozen=True): - """Normalized sparse-attention indexer description for one model.""" + """Compressed sparse-attention indexer configuration.""" - variant: str # QSA_VARIANT_COMPRESSED | QSA_VARIANT_TOKENWISE n_heads: int # index query heads kv_heads: int # index key/value heads head_dim: int # per-head index dimension budget: int # tokens selected per query row - compress_ratio: int # 1 for tokenwise variants - rope_mode: str # rotary layout the indexer expects + compress_ratio: int @property def block_topk(self) -> int: - """Compressed blocks selected per query row (== budget for tokenwise).""" + """Compressed blocks selected per query row.""" return self.budget // self.compress_ratio @@ -64,11 +37,6 @@ def _text_config(config): return getattr(config, "text_config", config) -def _is_qwen_family(config) -> bool: - model_type = str(getattr(config, "model_type", "") or "") - return model_type.startswith("qwen") - - def _require_fields(config, fields) -> dict: missing = [name for name in fields if getattr(config, name, None) is None] if missing: @@ -99,40 +67,11 @@ def _parse_compressed(text_config) -> QSAProfile: f"to be one of {sorted(_COMPRESSED_BLOCK_TOPK)}, got {budget // ratio}" ) return QSAProfile( - variant=QSA_VARIANT_COMPRESSED, n_heads=values["indexer_n_heads"], kv_heads=values["indexer_kv_heads"], head_dim=values["indexer_head_dim"], budget=budget, compress_ratio=ratio, - # The compressed indexer consumes the Qwen4-Exp layer's own (m)rope. - rope_mode=QSA_ROPE_MROPE, - ) - - -def _parse_tokenwise(text_config) -> QSAProfile: - values = _require_fields(text_config, _TOKENWISE_FIELDS) - if any(value <= 0 for value in values.values()): - raise ValueError(f"QSA config values must be positive: {values}") - if values["index_topk"] != _TOKENWISE_BUDGET: - raise ValueError( - f"fast_topk_v2 only supports index_topk = {_TOKENWISE_BUDGET}, " - f"got {values['index_topk']}" - ) - if values["index_kv_heads"] != 1: - raise ValueError( - f"QSA tokenwise index requires index_kv_heads = 1 (MQA), " - f"got {values['index_kv_heads']}" - ) - return QSAProfile( - variant=QSA_VARIANT_TOKENWISE, - n_heads=values["index_n_heads"], - kv_heads=values["index_kv_heads"], - head_dim=values["index_head_dim"], - budget=values["index_topk"], - compress_ratio=1, - # The tokenwise indexer owns plain per-token rotary positions. - rope_mode=QSA_ROPE_PLAIN, ) @@ -144,34 +83,19 @@ def parse_qsa_profile(config) -> Optional[QSAProfile]: text_config = _text_config(config) if text_config is None: return None - has_compressed = getattr(text_config, "indexer_n_heads", None) is not None - has_tokenwise = getattr( - text_config, "index_topk", None - ) is not None and _is_qwen_family(text_config) - if has_compressed and has_tokenwise: - raise ValueError( - "Ambiguous QSA config: both compressed (indexer_*) and tokenwise " - "(index_*) indexer fields are set" - ) - if has_compressed: + if getattr(text_config, "indexer_n_heads", None) is not None: return _parse_compressed(text_config) - if has_tokenwise: - return _parse_tokenwise(text_config) return None def is_qwen_qsa(config) -> bool: - """Return whether the config describes a supported Qwen QSA variant.""" + """Return whether the config describes Qwen compressed QSA.""" return parse_qsa_profile(config) is not None __all__ = [ "QSAProfile", - "QSA_ROPE_MROPE", - "QSA_ROPE_PLAIN", - "QSA_VARIANT_COMPRESSED", - "QSA_VARIANT_TOKENWISE", "is_qwen_qsa", "parse_qsa_profile", ] diff --git a/python/sglang/srt/layers/attention/qsa/dsa_indexer.py b/python/sglang/srt/layers/attention/qsa/dsa_indexer.py deleted file mode 100644 index 88f84cb86..000000000 --- a/python/sglang/srt/layers/attention/qsa/dsa_indexer.py +++ /dev/null @@ -1,337 +0,0 @@ -"""Tokenwise (per-token) QSA indexer for Qwen3Next-DSA models. - -A tokenwise profile has ``compress_ratio = 1`` and ``block_topk = budget = 2048``; -it never consumes the compressed-only MQA inputs. -Only the BF16 torch reference path is implemented; -requesting the FP8 or TileLang fast paths fails loudly. -""" - -from __future__ import annotations - -import logging - -import torch - -from sglang.srt.layers.attention.qsa.config import ( - QSA_VARIANT_TOKENWISE, - parse_qsa_profile, -) -from sglang.srt.layers.attention.qsa.kernel import qsa_fast_topk -from sglang.srt.layers.attention.qsa.qsa_indexer import _qsa_prefill_row_chunk_size -from sglang.srt.layers.layernorm import GemmaRMSNorm -from sglang.srt.layers.linear import ReplicatedLinear -from sglang.srt.layers.rotary_embedding import get_rope_wrapper -from sglang.srt.layers.utils import MultiPlatformOp - -logger = logging.getLogger(__name__) - - -def torch_dsa_weighted_mqa_logits( - q: torch.Tensor, - w: torch.Tensor, - k: torch.Tensor, - score_scale: float, -) -> torch.Tensor: - """Lightning-Index scoring reference: ReLU dot-product weighted per head.""" - - if k.ndim == 4: - if k.shape[2] != 1 or k.shape[0] != q.shape[0]: - raise ValueError( - "tokenwise MQA requires per-row k [rows, keys, 1, head_dim], " - f"got {k.shape}" - ) - scores = torch.relu(torch.einsum("mhd,mkhd->mkh", q.float(), k.float())) - else: - if k.ndim != 3 or k.shape[1] != 1: - raise ValueError( - f"tokenwise MQA requires k [keys, 1, head_dim], got {k.shape}" - ) - scores = torch.relu(torch.einsum("mhd,khd->mkh", q.float(), k.float())) - return (scores * w.float().unsqueeze(1)).sum(dim=-1) / score_scale - - -class QwenDSAIndexer(MultiPlatformOp): - """Tokenwise Lightning Indexer with the compressed ``QSAIndexer`` forward contract; - returns per-row logical token indices consumed as ``topk_indices``.""" - - def __init__( - self, - config, - layer_id: int, - quant_config=None, - prefix: str = "", - page_size: int = 64, - max_model_len=None, - ) -> None: - super().__init__() - profile = parse_qsa_profile(config) - if profile is None or profile.variant != QSA_VARIANT_TOKENWISE: - raise ValueError( - "QwenDSAIndexer requires a tokenwise QSA config (index_topk/), " - f"got profile={profile}" - ) - if page_size != 64: - # The paged index-K layout and every fast path assume 64-token - # pages, matching qsa_0511. - raise ValueError(f"tokenwise QSA requires page_size = 64, got {page_size}") - self.qsa_profile = profile - self.layer_id = int(layer_id) - self.index_n_heads = profile.n_heads - self.index_kv_heads = profile.kv_heads - self.index_head_dim = profile.head_dim - self.token_topk = profile.budget - self.score_scale = float(profile.head_dim) ** 0.5 - self.page_size = page_size - self.max_model_len = max_model_len - - # Fused Q/K/W projection. Output layout: - # q_raw: [M, index_n_heads * index_head_dim] - # k_raw: [M, index_kv_heads * index_head_dim] - # w: [M, index_n_heads] per-head scalar weight - self.index_q_dim = self.index_n_heads * self.index_head_dim - self.index_k_dim = self.index_kv_heads * self.index_head_dim - self.index_w_dim = self.index_n_heads - self.index_qkw_proj = ReplicatedLinear( - config.hidden_size, - self.index_q_dim + self.index_k_dim + self.index_w_dim, - bias=False, - quant_config=quant_config, - prefix=f"{prefix}.index_qkw_proj" if prefix else "index_qkw_proj", - ) - self.index_q_layernorm = GemmaRMSNorm( - self.index_head_dim, eps=getattr(config, "rms_norm_eps", 1e-6) - ) - self.index_k_layernorm = GemmaRMSNorm( - self.index_head_dim, eps=getattr(config, "rms_norm_eps", 1e-6) - ) - - # The indexer keeps its own RoPE instance shaped for index_head_dim; - # its rotary width follows the main attention's partial_rotary_factor. - rope_scaling = getattr(config, "rope_scaling", None) - if rope_scaling is None: - rope_scaling = getattr(config, "rope_parameters", None) - rope_theta = getattr(config, "rope_theta", 10000) - if isinstance(rope_scaling, dict) and "rope_theta" in rope_scaling: - rope_theta = rope_scaling["rope_theta"] - main_head_dim = getattr(config, "head_dim", None) - if main_head_dim is None: - main_head_dim = getattr(config, "hidden_size") // getattr( - config, "num_attention_heads" - ) - partial_rotary_factor = getattr(config, "partial_rotary_factor", None) - if partial_rotary_factor is None and isinstance(rope_scaling, dict): - partial_rotary_factor = rope_scaling.get("partial_rotary_factor") - if partial_rotary_factor is None: - partial_rotary_factor = 1.0 - indexer_rotary_dim = min( - self.index_head_dim, int(main_head_dim * float(partial_rotary_factor)) - ) - if indexer_rotary_dim <= 0 or indexer_rotary_dim % 2 != 0: - raise ValueError( - "tokenwise QSA indexer requires a positive even rotary dim, got " - f"{indexer_rotary_dim=} from {main_head_dim=} and " - f"{partial_rotary_factor=}" - ) - self.rotary_emb = get_rope_wrapper( - head_size=self.index_head_dim, - rotary_dim=indexer_rotary_dim, - max_position=getattr(config, "max_position_embeddings", 8192), - base=rope_theta, - rope_scaling=rope_scaling if isinstance(rope_scaling, dict) else None, - is_neox_style=True, - dtype=torch.get_default_dtype(), - ) - - def project_qkw(self, hidden_states: torch.Tensor, positions: torch.Tensor): - """Fused Q/K/W projection, per-head RMS norm and indexer RoPE.""" - - qkw, _ = self.index_qkw_proj(hidden_states) - q_raw, k_raw, w = torch.split( - qkw, [self.index_q_dim, self.index_k_dim, self.index_w_dim], dim=-1 - ) - q = self.index_q_layernorm(q_raw.reshape(-1, self.index_head_dim)).reshape( - -1, self.index_n_heads, self.index_head_dim - ) - k = self.index_k_layernorm(k_raw.reshape(-1, self.index_head_dim)).reshape( - -1, self.index_kv_heads, self.index_head_dim - ) - q, k = self.rotary_emb(positions, q, k) - return q, w, k - - def forward_cuda( - self, - hidden_states: torch.Tensor, - positions: torch.Tensor, - forward_batch, - indexer_metadata, - ) -> torch.Tensor: - forward_mode = forward_batch.forward_mode - is_target_verify = getattr(forward_mode, "is_target_verify", lambda: False)() - is_draft_extend = getattr(forward_mode, "is_draft_extend", lambda **_: False)( - include_v2=True - ) - is_paged = forward_mode.is_decode() or is_target_verify or is_draft_extend - if is_paged: - # Paged rows take their causal length from the paged metadata, - # not the model's RoPE coordinate, as in the compressed QSAIndexer. - logical_positions = indexer_metadata.get_seqlens_expanded() - 1 - else: - logical_positions = getattr(forward_batch, "positions", None) - if logical_positions is None: - logical_positions = positions[0] if positions.ndim == 2 else positions - logical_positions = logical_positions.flatten() - - # DP padding adds token rows that belong to no request; - # token_to_batch_idx is the source of truth for semantic rows. - num_valid_tokens = indexer_metadata.get_token_to_batch_idx().numel() - if logical_positions.numel() < num_valid_tokens: - raise ValueError( - "tokenwise QSA logical positions are shorter than the request " - f"mapping: positions={logical_positions.numel()}, " - f"mapping={num_valid_tokens}" - ) - if hidden_states.shape[0] < num_valid_tokens: - raise ValueError( - "tokenwise QSA hidden states are shorter than the request " - f"mapping: hidden={hidden_states.shape[0]}, " - f"mapping={num_valid_tokens}" - ) - position_tokens = ( - positions.shape[-1] if positions.ndim == 2 else positions.numel() - ) - if position_tokens < num_valid_tokens: - raise ValueError( - "tokenwise QSA RoPE positions are shorter than the request " - f"mapping: positions={position_tokens}, " - f"mapping={num_valid_tokens}" - ) - - logical_positions = logical_positions[:num_valid_tokens] - hidden_states = hidden_states[:num_valid_tokens] - positions = ( - positions[:, :num_valid_tokens] - if positions.ndim == 2 - else positions[:num_valid_tokens] - ) - if num_valid_tokens == 0: - return torch.empty( - (0, self.token_topk), - dtype=torch.int32, - device=hidden_states.device, - ) - - q, w, k = self.project_qkw(hidden_states, positions) - - pool = indexer_metadata.token_to_kv_pool - out_cache_loc = getattr(indexer_metadata, "out_cache_loc", None) - if out_cache_loc is None: - out_cache_loc = forward_batch.out_cache_loc - pool.set_dsa_index_k_buffer(self.layer_id, out_cache_loc[:num_valid_tokens], k) - - if is_paged: - return self._select_paged(q, w, indexer_metadata) - return self._select_prefill(q, w, logical_positions, indexer_metadata) - - def _select_paged( - self, - q: torch.Tensor, - w: torch.Tensor, - indexer_metadata, - ) -> torch.Tensor: - """Per-query-row top-k over ``[0, row_len)`` for paged modes.""" - - pool = indexer_metadata.token_to_kv_pool - index_k = pool.get_dsa_index_k_buffer(self.layer_id) - sequence_lengths = indexer_metadata.sequence_lengths.to(torch.int32) - table = indexer_metadata.token_slot_table - rows, max_len = table.shape - if rows != indexer_metadata.token_to_batch_idx.numel(): - raise ValueError( - "tokenwise QSA paged modes need one slot-table row per query " - f"row: table_rows={rows}, " - f"mapping={indexer_metadata.token_to_batch_idx.numel()}" - ) - output = torch.full( - (rows, self.token_topk), -1, dtype=torch.int32, device=q.device - ) - row_chunk = _qsa_prefill_row_chunk_size(rows, max_len, self.index_n_heads) - table_long = table.long() - for row_start in range(0, rows, row_chunk): - row_end = min(row_start + row_chunk, rows) - # Table columns at/after each row's length hold stale slots; the - # gathers stay in range and fast_topk masks them out by length. - k_chunk = index_k.index_select(0, table_long[row_start:row_end].reshape(-1)) - k_chunk = k_chunk.reshape(row_end - row_start, max_len, 1, -1) - logits = torch_dsa_weighted_mqa_logits( - q[row_start:row_end], - w[row_start:row_end], - k_chunk, - self.score_scale, - ) - lengths = sequence_lengths[row_start:row_end] - selected = qsa_fast_topk( - logits, - torch.zeros_like(lengths), - lengths.clamp(min=0, max=max_len), - topk=self.token_topk, - ) - output[row_start:row_end].copy_(selected) - return output - - def _select_prefill( - self, - q: torch.Tensor, - w: torch.Tensor, - logical_positions: torch.Tensor, - indexer_metadata, - ) -> torch.Tensor: - """Packed per-sequence top-k with causal windows for extend modes.""" - - pool = indexer_metadata.token_to_kv_pool - index_k = pool.get_dsa_index_k_buffer(self.layer_id) - sequence_lengths = indexer_metadata.sequence_lengths.to(torch.int32) - table = indexer_metadata.token_slot_table - query_sequence_ids = indexer_metadata.token_to_batch_idx.long() - row_ends_all = (logical_positions.to(torch.int32) + 1).clamp( - min=0, max=table.shape[1] - ) - rows = q.shape[0] - output = torch.full( - (rows, self.token_topk), -1, dtype=torch.int32, device=q.device - ) - for sequence_id in range(sequence_lengths.numel()): - seq_len = int(sequence_lengths[sequence_id].item()) - row_mask = query_sequence_ids == sequence_id - if seq_len <= 0 or not bool(row_mask.any()): - continue - row_indices = row_mask.nonzero(as_tuple=True)[0] - slots = table[sequence_id, :seq_len].long() - k_seq = index_k.index_select(0, slots) - row_chunk = _qsa_prefill_row_chunk_size( - row_indices.numel(), seq_len, self.index_n_heads - ) - for chunk_start in range(0, row_indices.numel(), row_chunk): - chunk_rows = row_indices[chunk_start : chunk_start + row_chunk] - row_ends = row_ends_all.index_select(0, chunk_rows) - logits = torch_dsa_weighted_mqa_logits( - q.index_select(0, chunk_rows), - w.index_select(0, chunk_rows), - k_seq, - self.score_scale, - ) - selected = qsa_fast_topk( - logits, - torch.zeros_like(row_ends), - row_ends, - topk=self.token_topk, - ) - # Tensor indexing returns a copy on read; use index_put style - # assignment or the selection would never reach `output`. - output[chunk_rows] = selected - return output - - -__all__ = [ - "QwenDSAIndexer", - "torch_dsa_weighted_mqa_logits", -] diff --git a/python/sglang/srt/layers/attention/qsa/glue.py b/python/sglang/srt/layers/attention/qsa/glue.py index 326895770..b824d8a72 100644 --- a/python/sglang/srt/layers/attention/qsa/glue.py +++ b/python/sglang/srt/layers/attention/qsa/glue.py @@ -2,10 +2,7 @@ from __future__ import annotations -from sglang.srt.layers.attention.qsa.config import ( - QSA_VARIANT_COMPRESSED, - parse_qsa_profile, -) +from sglang.srt.layers.attention.qsa.config import parse_qsa_profile def build_qsa_indexer( @@ -22,27 +19,14 @@ def build_qsa_indexer( raise ValueError( "build_qsa_indexer requires a config with a QSA indexer schema" ) - if profile.variant == QSA_VARIANT_COMPRESSED: - # The compressed indexer reuses the layer's own Qwen4-Exp RoPE - # (mrope); there is intentionally no plain-rope path for it here. - from sglang.srt.layers.attention.qsa.qsa_indexer import QSAIndexer + from sglang.srt.layers.attention.qsa.qsa_indexer import QSAIndexer - return QSAIndexer( - config=config, - layer_id=layer_id, - quant_config=quant_config, - prefix=prefix, - rotary_emb=rotary_emb, - ) - # Tokenwise (Qwen3Next-DSA): the Lightning Indexer owns its plain - # per-token RoPE; a shared layer rotary is neither needed nor accepted. - from sglang.srt.layers.attention.qsa.dsa_indexer import QwenDSAIndexer - - return QwenDSAIndexer( + return QSAIndexer( config=config, layer_id=layer_id, quant_config=quant_config, prefix=prefix, + rotary_emb=rotary_emb, ) diff --git a/python/sglang/srt/layers/attention/qwen_sparse_attn_backend.py b/python/sglang/srt/layers/attention/qwen_sparse_attn_backend.py index d42d6695b..9cd60707f 100644 --- a/python/sglang/srt/layers/attention/qwen_sparse_attn_backend.py +++ b/python/sglang/srt/layers/attention/qwen_sparse_attn_backend.py @@ -18,7 +18,6 @@ import torch.nn.functional as F from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.qsa.config import ( - QSA_VARIANT_COMPRESSED, is_qwen_qsa, parse_qsa_profile, ) @@ -183,7 +182,6 @@ class QwenSparseAttnBackend(AttentionBackend): config = getattr(model_config, "hf_text_config", None) if config is None: config = getattr(model_config, "hf_config", None) - # Compressed (Qwen4-Exp) and tokenwise (Qwen3Next-DSA) QSA share this backend. self.qsa_profile = parse_qsa_profile(config) self.max_context_len = int(getattr(model_config, "context_len", 0)) self.compress_ratio = ( @@ -631,77 +629,69 @@ class QwenSparseAttnBackend(AttentionBackend): f"mapping={token_to_batch_idx.numel()}, " f"positions={num_position_tokens}" ) - write_locs = None - group_positions = None - group_sequence_ids = None - group_member_rows = None decode_page_table = None decode_lengths = None decode_logical_positions = None pending_ring_slots = None compress_group_ring_locs = None extend_rope_matrix = None - if ( - self.qsa_profile is None - or self.qsa_profile.variant == QSA_VARIANT_COMPRESSED - ): - write_locs, group_positions, group_sequence_ids, group_member_rows = ( - self._qsa_build_write_plan( - forward_batch=forward_batch, - speculative_paged=speculative_paged, - token_slot_table=token_slot_table, - sequence_lengths=sequence_lengths, - ) + write_locs, group_positions, group_sequence_ids, group_member_rows = ( + self._qsa_build_write_plan( + forward_batch=forward_batch, + speculative_paged=speculative_paged, + token_slot_table=token_slot_table, + sequence_lengths=sequence_lengths, ) - decode_like = speculative_paged or forward_batch.forward_mode.is_decode() + ) + decode_like = speculative_paged or forward_batch.forward_mode.is_decode() + if decode_like: + decode_logical_positions = ( + logical_positions.to(torch.int32) + if speculative_paged + else sequence_lengths - 1 + ) + ring_logical_positions = decode_logical_positions + else: + extend_positions = forward_batch.positions + if extend_positions.ndim == 2: + extend_positions = extend_positions[0] + ring_logical_positions = extend_positions.flatten()[ + : token_to_batch_idx.numel() + ] + if not self.should_reuse_mtp_sparse_indices(forward_batch): if decode_like: - decode_logical_positions = ( - logical_positions.to(torch.int32) - if speculative_paged - else sequence_lengths - 1 - ) - ring_logical_positions = decode_logical_positions - else: - extend_positions = forward_batch.positions - if extend_positions.ndim == 2: - extend_positions = extend_positions[0] - ring_logical_positions = extend_positions.flatten()[ - : token_to_batch_idx.numel() - ] - if not self.should_reuse_mtp_sparse_indices(forward_batch): - if decode_like: - pool = self.token_to_kv_pool - decode_page_table, decode_lengths = compressed_decode_view( - compressed_page_size=pool.qsa_compressed_page_size, - compress_ratio=pool.qsa_compress_ratio, - sequence_lengths=sequence_lengths, - token_slot_table=token_slot_table, - ) - pending_ring_slots = build_pending_ring_slots( - token_to_batch_idx=token_to_batch_idx, - req_pool_indices=row_req_pool_indices, + pool = self.token_to_kv_pool + decode_page_table, decode_lengths = compressed_decode_view( + compressed_page_size=pool.qsa_compressed_page_size, + compress_ratio=pool.qsa_compress_ratio, sequence_lengths=sequence_lengths, - logical_positions=ring_logical_positions, - compress_ratio=self.compress_ratio, - is_extend=group_member_rows is not None, + token_slot_table=token_slot_table, ) - if write_locs.numel(): - if group_member_rows is not None: - rope_source = ( - forward_batch.mrope_positions - if forward_batch.mrope_positions is not None - else forward_batch.positions - ) - extend_rope_matrix = build_rope_position_matrix( - rope_source, token_to_batch_idx.numel() - ) - else: - compress_group_ring_locs = build_group_ring_slots( - req_pool_indices=row_req_pool_indices, - group_end_positions=group_positions.long(), - sequence_ids=group_sequence_ids.long(), - compress_ratio=self.compress_ratio, - ) + pending_ring_slots = build_pending_ring_slots( + token_to_batch_idx=token_to_batch_idx, + req_pool_indices=row_req_pool_indices, + sequence_lengths=sequence_lengths, + logical_positions=ring_logical_positions, + compress_ratio=self.compress_ratio, + is_extend=group_member_rows is not None, + ) + if write_locs.numel(): + if group_member_rows is not None: + rope_source = ( + forward_batch.mrope_positions + if forward_batch.mrope_positions is not None + else forward_batch.positions + ) + extend_rope_matrix = build_rope_position_matrix( + rope_source, token_to_batch_idx.numel() + ) + else: + compress_group_ring_locs = build_group_ring_slots( + req_pool_indices=row_req_pool_indices, + group_end_positions=group_positions.long(), + sequence_ids=group_sequence_ids.long(), + compress_ratio=self.compress_ratio, + ) indexer_metadata = QSAIndexerMetadata( sequence_lengths=sequence_lengths, token_to_batch_idx=token_to_batch_idx, @@ -835,17 +825,6 @@ class QwenSparseAttnBackend(AttentionBackend): ] self._extend_lens_pin_idx = 0 - def _require_compressed_cuda_graph_support(self) -> None: - if ( - self.qsa_profile is not None - and self.qsa_profile.variant != QSA_VARIANT_COMPRESSED - ): - raise NotImplementedError( - "QSA tokenwise CUDA-graph execution requires graph-stable " - "indexer metadata, which is not available in this tree yet; " - "run tokenwise QSA with --disable-cuda-graph" - ) - def _capture_cuda_graph_metadata( self, *, @@ -856,7 +835,6 @@ class QwenSparseAttnBackend(AttentionBackend): forward_mode, spec_info, ) -> None: - self._require_compressed_cuda_graph_support() self._require_chain_speculation(forward_mode, spec_info) if self.token_to_kv_pool is None: self.token_to_kv_pool = getattr(self.runner, "token_to_kv_pool", None) diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index 82784f11e..8c0203e61 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -1887,25 +1887,16 @@ class KVCacheConfigurator: else mha_pool_class ) from sglang.srt.layers.attention.qsa.config import ( - QSA_VARIANT_TOKENWISE, parse_qsa_profile, ) from sglang.srt.mem_cache.qsa_kv_pool import ( QSATokenToKVPool, - QwenDSATokenToKVPool, ) qsa_profile = parse_qsa_profile(self.model_config.hf_config) if qsa_profile is None: pool_class = HybridLinearKVPool extra_args["use_mla"] = self.use_mla_backend - elif qsa_profile.variant == QSA_VARIANT_TOKENWISE: - pool_class = QwenDSATokenToKVPool - extra_args.update( - qsa_index_kv_heads=qsa_profile.kv_heads, - qsa_index_head_dim=qsa_profile.head_dim, - qsa_token_budget=qsa_profile.budget, - ) else: pool_class = QSATokenToKVPool extra_args.update( diff --git a/python/sglang/srt/mem_cache/qsa_kv_pool.py b/python/sglang/srt/mem_cache/qsa_kv_pool.py index 49d69bdb7..61308c2e3 100644 --- a/python/sglang/srt/mem_cache/qsa_kv_pool.py +++ b/python/sglang/srt/mem_cache/qsa_kv_pool.py @@ -283,107 +283,3 @@ class QSATokenToKVPool(HybridLinearKVPool): + self.qsa_rope_position_buffer.numel() * 8 ) return k_size + qsa_k_size, v_size - - -class QwenDSATokenToKVPool(HybridLinearKVPool): - """Hybrid KV pool carrying the per-token index-K cache of tokenwise QSA: - a ``[size + page_size, index_kv_heads, index_head_dim]`` BF16 buffer per DSA layer, - addressed by raw KV slots; the FP8 deep_gemm layout is deliberately absent.""" - - index_state_dtype = torch.bfloat16 - - @classmethod - def qsa_bytes_per_token( - cls, *, kv_heads: int, head_dim: int, num_layers: int - ) -> int: - return ( - _index_k_bytes( - kv_heads=kv_heads, head_dim=head_dim, dtype=cls.index_state_dtype - ) - * num_layers - ) - - def __init__( - self, - *, - size: int, - dtype: torch.dtype, - page_size: int, - head_num: int, - head_dim: int, - full_attention_layer_ids: List[int], - device: str, - mamba_pool: MambaPool, - qsa_index_kv_heads: int, - qsa_index_head_dim: int, - qsa_token_budget: int, - enable_memory_saver: bool = False, - enable_kv_cache_copy: bool = False, - start_layer: Optional[int] = None, - full_kv_pool_class: Optional[type] = None, - quant_method=None, - post_capture_active: bool = False, - ): - if page_size != 64: - raise ValueError( - "tokenwise QSA requires KV-cache page_size 64 for its paged " - f"indexer buffer, got {page_size}" - ) - self.dsa_index_k_buffer_pool = [] - super().__init__( - size=size, - dtype=dtype, - page_size=page_size, - head_num=head_num, - head_dim=head_dim, - full_attention_layer_ids=full_attention_layer_ids, - device=device, - mamba_pool=mamba_pool, - enable_memory_saver=enable_memory_saver, - enable_kv_cache_copy=enable_kv_cache_copy, - use_mla=False, - start_layer=start_layer, - full_kv_pool_class=full_kv_pool_class, - quant_method=quant_method, - post_capture_active=post_capture_active, - ) - if qsa_index_kv_heads != 1: - raise ValueError( - f"tokenwise QSA requires index_kv_heads = 1 (MQA), got " - f"{qsa_index_kv_heads}" - ) - if min(qsa_index_kv_heads, qsa_index_head_dim, qsa_token_budget) <= 0: - raise ValueError("QSA cache configuration values must be positive") - self.qsa_compress_ratio = 1 - self.qsa_index_kv_heads = int(qsa_index_kv_heads) - self.qsa_index_head_dim = int(qsa_index_head_dim) - self.qsa_token_topk = int(qsa_token_budget) - self.qsa_block_topk = int(qsa_token_budget) - state_size = size + page_size - self.dsa_index_k_buffer_pool = [ - torch.zeros( - (state_size, self.qsa_index_kv_heads, self.qsa_index_head_dim), - dtype=self.index_state_dtype, - device=device, - ) - for _ in full_attention_layer_ids - ] - k_size, v_size = self.get_kv_size_bytes() - self.mem_usage = (k_size + v_size) / GB - - def set_dsa_index_k_buffer( - self, layer_id: int, loc: torch.Tensor, index_k: torch.Tensor - ) -> None: - buffer = self.get_dsa_index_k_buffer(layer_id) - buffer[loc.long()] = index_k.to(buffer.dtype) - - def get_dsa_index_k_buffer(self, layer_id: int) -> torch.Tensor: - return self.dsa_index_k_buffer_pool[self._transfer_full_attention_id(layer_id)] - - def get_kv_size_bytes(self): - k_size, v_size = super().get_kv_size_bytes() - dsa_k_size = sum( - tensor.numel() * tensor.element_size() - for tensor in self.dsa_index_k_buffer_pool - ) - return k_size + dsa_k_size, v_size diff --git a/python/sglang/srt/model_executor/pool_configurator.py b/python/sglang/srt/model_executor/pool_configurator.py index d3cb42449..8ec857a6a 100644 --- a/python/sglang/srt/model_executor/pool_configurator.py +++ b/python/sglang/srt/model_executor/pool_configurator.py @@ -431,12 +431,10 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): @staticmethod def _compute_qsa_cell_size(*, hf_config, num_layers: int) -> int: from sglang.srt.layers.attention.qsa.config import ( - QSA_VARIANT_COMPRESSED, parse_qsa_profile, ) from sglang.srt.mem_cache.qsa_kv_pool import ( QSATokenToKVPool, - QwenDSATokenToKVPool, ) if num_layers == 0: @@ -444,16 +442,10 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): qsa_profile = parse_qsa_profile(hf_config) if qsa_profile is None: return 0 - if qsa_profile.variant == QSA_VARIANT_COMPRESSED: - return QSATokenToKVPool.qsa_bytes_per_token( - kv_heads=qsa_profile.kv_heads, - head_dim=qsa_profile.head_dim, - compress_ratio=qsa_profile.compress_ratio, - num_layers=num_layers, - ) - return QwenDSATokenToKVPool.qsa_bytes_per_token( + return QSATokenToKVPool.qsa_bytes_per_token( kv_heads=qsa_profile.kv_heads, head_dim=qsa_profile.head_dim, + compress_ratio=qsa_profile.compress_ratio, num_layers=num_layers, ) diff --git a/python/sglang/srt/speculative/draft_utils.py b/python/sglang/srt/speculative/draft_utils.py index b7ff407fd..5c1306688 100644 --- a/python/sglang/srt/speculative/draft_utils.py +++ b/python/sglang/srt/speculative/draft_utils.py @@ -1,6 +1,6 @@ from typing import Optional -from sglang.srt.layers.attention.qsa.config import QSA_VARIANT_COMPRESSED, QSAProfile +from sglang.srt.layers.attention.qsa.config import QSAProfile from sglang.srt.runtime_context import attention_backends, get_spec from sglang.srt.utils.common import ( cpu_has_amx_support, @@ -175,10 +175,6 @@ class DraftBackendFactory: backend.decode_attention_backend_str = "qsa" def _create_qwen_qsa_draft_extend_backend(self): - if self.qsa_profile.variant != QSA_VARIANT_COMPRESSED: - # Tokenwise QSA has no graph-stable indexer metadata: draft extend - # stays eager instead of falling back to a dense backend. - return None from sglang.srt.layers.attention.qwen_sparse_attn_backend import ( QwenSparseAttnBackend, ) diff --git a/test/registered/kernel/qsa/test_qsa.py b/test/registered/kernel/qsa/test_qsa.py index be5e46a24..8bb75e625 100644 --- a/test/registered/kernel/qsa/test_qsa.py +++ b/test/registered/kernel/qsa/test_qsa.py @@ -7,7 +7,6 @@ import torch from sglang.kernels.ops.attention import qwen38_qsa_sm121_varlen from sglang.srt.configs.qwen4_exp import Qwen4ExpConfig from sglang.srt.layers.attention import qwen_sparse_attn_backend as qsa_backend_module -from sglang.srt.layers.attention.qsa import dsa_indexer as dsa_indexer_module from sglang.srt.layers.attention.qsa import qsa_indexer as qsa_indexer_module from sglang.srt.layers.attention.qsa.kernel import ( expand_qsa_block_indices, @@ -261,22 +260,8 @@ def _compressed_config_namespace(**overrides): return SimpleNamespace(**fields) -def _tokenwise_config_namespace(**overrides): - fields = dict( - model_type="qwen3_5", - index_topk=2048, - index_n_heads=64, - index_kv_heads=1, - index_head_dim=128, - ) - fields.update(overrides) - return SimpleNamespace(**fields) - - def test_qsa_profile_parses_compressed_qwen4_exp_schema(): from sglang.srt.layers.attention.qsa.config import ( - QSA_ROPE_MROPE, - QSA_VARIANT_COMPRESSED, is_qwen_qsa, parse_qsa_profile, ) @@ -297,14 +282,12 @@ def test_qsa_profile_parses_compressed_qwen4_exp_schema(): ) for config in (wrapped, _compressed_config_namespace()): profile = parse_qsa_profile(config) - assert profile.variant == QSA_VARIANT_COMPRESSED assert profile.n_heads == 8 assert profile.kv_heads == 1 assert profile.head_dim == 128 assert profile.budget == TOKEN_TOPK assert profile.compress_ratio == COMPRESS_RATIO assert profile.block_topk == BLOCK_TOPK - assert profile.rope_mode == QSA_ROPE_MROPE assert is_qwen_qsa(config) # The legacy backend module keeps re-exporting the shared detector. assert qsa_backend_module.is_qwen_qsa is is_qwen_qsa @@ -332,7 +315,7 @@ def test_qsa_profile_rejects_malformed_compressed_schema(): raise AssertionError(f"{name} compressed config must be rejected") -def test_qsa_glue_builds_indexer_per_variant(monkeypatch): +def test_qsa_glue_builds_compressed_indexer(monkeypatch): from sglang.srt.layers.attention.qsa.glue import build_qsa_indexer recorded = {} @@ -360,24 +343,6 @@ def test_qsa_glue_builds_indexer_per_variant(monkeypatch): config=config, layer_id=7, quant_config="qc", prefix="p", rotary_emb=rotary ) - # Tokenwise configs build the Lightning Indexer through the same glue. - class _FakeDSAIndexer: - def __init__(self, config, layer_id, quant_config=None, prefix="", **kw): - recorded.update( - dsa_config=config, - dsa_layer_id=layer_id, - dsa_quant_config=quant_config, - dsa_prefix=prefix, - ) - - monkeypatch.setattr(dsa_indexer_module, "QwenDSAIndexer", _FakeDSAIndexer) - dsa_indexer = build_qsa_indexer( - _tokenwise_config_namespace(), layer_id=2, prefix="q" - ) - assert isinstance(dsa_indexer, _FakeDSAIndexer) - assert recorded["dsa_layer_id"] == 2 - assert recorded["dsa_prefix"] == "q" - try: build_qsa_indexer(SimpleNamespace(), layer_id=0, rotary_emb=rotary) except ValueError as exc: @@ -438,9 +403,6 @@ def test_qsa_draft_extend_backend_decision_follows_profile(): assert isinstance(backend, QwenSparseAttnBackend) assert backend.runner is compressed.draft_model_runner assert backend.decode_attention_backend_str == "qsa" - # Tokenwise profiles stay eager (no graph-stable indexer metadata); they - # must never silently fall back to a dense backend either. - assert factory(_tokenwise_config_namespace()).create_draft_extend_backend() is None def _make_mtp_draft_batch(steps: int, seq_lens=(8, 16), loc_base: int = 40):