[Qwen 3.8 Next] Remove unused tokenwise QSA implementation and tests (#38960)
This commit is contained in:
@@ -39,7 +39,7 @@ SGL_DEVICE auto convert_to_uint32(float x) -> uint32_t {
|
|||||||
|
|
||||||
// When length <= kTopK, write the indices directly.
|
// When length <= kTopK, write the indices directly.
|
||||||
template <int kTopK>
|
template <int kTopK>
|
||||||
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;
|
const auto tid = threadIdx.x;
|
||||||
for (int i = tid; i < kTopK; i += kThreadsPerBlock) {
|
for (int i = tid; i < kTopK; i += kThreadsPerBlock) {
|
||||||
indice[i] = (i < length) ? i : -1;
|
indice[i] = (i < length) ? i : -1;
|
||||||
@@ -217,12 +217,12 @@ __global__ __launch_bounds__(fast_topk_detail::kThreadsPerBlock) void fast_topk_
|
|||||||
device::PDLWaitPrimary<kUsePDL>();
|
device::PDLWaitPrimary<kUsePDL>();
|
||||||
|
|
||||||
const auto bid = static_cast<uint64_t>(blockIdx.x);
|
const auto bid = static_cast<uint64_t>(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 length = params.lengths[bid];
|
||||||
const auto indice = params.indices + bid * kTopK;
|
const auto indice = params.indices + bid * kTopK;
|
||||||
const auto score = params.input + bid * params.input_stride;
|
const auto score = params.input + bid * params.input_stride;
|
||||||
if (length <= kTopK) {
|
if (length <= kTopK) {
|
||||||
naive_topk<kTopK>(score, indice, length);
|
naive_topk<kTopK>(indice, length);
|
||||||
} else {
|
} else {
|
||||||
radix_select_topk<kTopK>(score, indice, row_start, length);
|
radix_select_topk<kTopK>(score, indice, row_start, length);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
overrides["page_size"] = 64 if sm100_default_attn_backend == "trtllm_mha" else 1
|
||||||
|
|
||||||
from sglang.srt.layers.attention.qsa.config import (
|
from sglang.srt.layers.attention.qsa.config import (
|
||||||
QSA_VARIANT_COMPRESSED,
|
|
||||||
parse_qsa_profile,
|
parse_qsa_profile,
|
||||||
)
|
)
|
||||||
|
|
||||||
profile = parse_qsa_profile(hf_config)
|
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.
|
# Compressed slot = full_slot // ratio; all backends need page-aligned pages.
|
||||||
# mamba_radix_cache_strategy resolves later, so do not gate on it.
|
# mamba_radix_cache_strategy resolves later, so do not gate on it.
|
||||||
overrides["page_size"] = 64
|
overrides["page_size"] = 64
|
||||||
|
|||||||
@@ -8,7 +8,6 @@ __all__ = [
|
|||||||
"QSAIndexer",
|
"QSAIndexer",
|
||||||
"QSAIndexerMetadata",
|
"QSAIndexerMetadata",
|
||||||
"QSAProfile",
|
"QSAProfile",
|
||||||
"QwenDSAIndexer",
|
|
||||||
"build_qsa_indexer",
|
"build_qsa_indexer",
|
||||||
"get_qsa_indexer_metadata",
|
"get_qsa_indexer_metadata",
|
||||||
"is_qwen_qsa",
|
"is_qwen_qsa",
|
||||||
@@ -21,10 +20,6 @@ def __getattr__(name):
|
|||||||
from sglang.srt.layers.attention.qsa.qsa_indexer import QSAIndexer
|
from sglang.srt.layers.attention.qsa.qsa_indexer import QSAIndexer
|
||||||
|
|
||||||
return QSAIndexer
|
return QSAIndexer
|
||||||
if name == "QwenDSAIndexer":
|
|
||||||
from sglang.srt.layers.attention.qsa.dsa_indexer import QwenDSAIndexer
|
|
||||||
|
|
||||||
return QwenDSAIndexer
|
|
||||||
if name == "QSAIndexerMetadata":
|
if name == "QSAIndexerMetadata":
|
||||||
from sglang.srt.layers.attention.qsa.metadata import QSAIndexerMetadata
|
from sglang.srt.layers.attention.qsa.metadata import QSAIndexerMetadata
|
||||||
|
|
||||||
|
|||||||
@@ -1,12 +1,4 @@
|
|||||||
"""Shared QSA profile parsing across model variants.
|
"""QSA profile parsing for Qwen4-Exp compressed indexing."""
|
||||||
|
|
||||||
``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``.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
@@ -14,14 +6,6 @@ from typing import Optional
|
|||||||
|
|
||||||
import msgspec
|
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 = (
|
_COMPRESSED_FIELDS = (
|
||||||
"indexer_n_heads",
|
"indexer_n_heads",
|
||||||
"indexer_kv_heads",
|
"indexer_kv_heads",
|
||||||
@@ -29,33 +13,22 @@ _COMPRESSED_FIELDS = (
|
|||||||
"indexer_budget",
|
"indexer_budget",
|
||||||
"indexer_compress_ratio",
|
"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.
|
# fast_topk_v2 only supports these compressed block top-k widths.
|
||||||
_COMPRESSED_BLOCK_TOPK = frozenset({512, 2048})
|
_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):
|
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
|
n_heads: int # index query heads
|
||||||
kv_heads: int # index key/value heads
|
kv_heads: int # index key/value heads
|
||||||
head_dim: int # per-head index dimension
|
head_dim: int # per-head index dimension
|
||||||
budget: int # tokens selected per query row
|
budget: int # tokens selected per query row
|
||||||
compress_ratio: int # 1 for tokenwise variants
|
compress_ratio: int
|
||||||
rope_mode: str # rotary layout the indexer expects
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def block_topk(self) -> int:
|
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
|
return self.budget // self.compress_ratio
|
||||||
|
|
||||||
@@ -64,11 +37,6 @@ def _text_config(config):
|
|||||||
return getattr(config, "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:
|
def _require_fields(config, fields) -> dict:
|
||||||
missing = [name for name in fields if getattr(config, name, None) is None]
|
missing = [name for name in fields if getattr(config, name, None) is None]
|
||||||
if missing:
|
if missing:
|
||||||
@@ -99,40 +67,11 @@ def _parse_compressed(text_config) -> QSAProfile:
|
|||||||
f"to be one of {sorted(_COMPRESSED_BLOCK_TOPK)}, got {budget // ratio}"
|
f"to be one of {sorted(_COMPRESSED_BLOCK_TOPK)}, got {budget // ratio}"
|
||||||
)
|
)
|
||||||
return QSAProfile(
|
return QSAProfile(
|
||||||
variant=QSA_VARIANT_COMPRESSED,
|
|
||||||
n_heads=values["indexer_n_heads"],
|
n_heads=values["indexer_n_heads"],
|
||||||
kv_heads=values["indexer_kv_heads"],
|
kv_heads=values["indexer_kv_heads"],
|
||||||
head_dim=values["indexer_head_dim"],
|
head_dim=values["indexer_head_dim"],
|
||||||
budget=budget,
|
budget=budget,
|
||||||
compress_ratio=ratio,
|
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)
|
text_config = _text_config(config)
|
||||||
if text_config is None:
|
if text_config is None:
|
||||||
return None
|
return None
|
||||||
has_compressed = getattr(text_config, "indexer_n_heads", None) is not None
|
if 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:
|
|
||||||
return _parse_compressed(text_config)
|
return _parse_compressed(text_config)
|
||||||
if has_tokenwise:
|
|
||||||
return _parse_tokenwise(text_config)
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def is_qwen_qsa(config) -> bool:
|
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
|
return parse_qsa_profile(config) is not None
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"QSAProfile",
|
"QSAProfile",
|
||||||
"QSA_ROPE_MROPE",
|
|
||||||
"QSA_ROPE_PLAIN",
|
|
||||||
"QSA_VARIANT_COMPRESSED",
|
|
||||||
"QSA_VARIANT_TOKENWISE",
|
|
||||||
"is_qwen_qsa",
|
"is_qwen_qsa",
|
||||||
"parse_qsa_profile",
|
"parse_qsa_profile",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -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",
|
|
||||||
]
|
|
||||||
@@ -2,10 +2,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from sglang.srt.layers.attention.qsa.config import (
|
from sglang.srt.layers.attention.qsa.config import parse_qsa_profile
|
||||||
QSA_VARIANT_COMPRESSED,
|
|
||||||
parse_qsa_profile,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def build_qsa_indexer(
|
def build_qsa_indexer(
|
||||||
@@ -22,9 +19,6 @@ def build_qsa_indexer(
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
"build_qsa_indexer requires a config with a QSA indexer schema"
|
"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(
|
return QSAIndexer(
|
||||||
@@ -34,16 +28,6 @@ def build_qsa_indexer(
|
|||||||
prefix=prefix,
|
prefix=prefix,
|
||||||
rotary_emb=rotary_emb,
|
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(
|
|
||||||
config=config,
|
|
||||||
layer_id=layer_id,
|
|
||||||
quant_config=quant_config,
|
|
||||||
prefix=prefix,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def resolve_qsa_sparse_backend(attn_backend):
|
def resolve_qsa_sparse_backend(attn_backend):
|
||||||
|
|||||||
@@ -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.base_attn_backend import AttentionBackend
|
||||||
from sglang.srt.layers.attention.qsa.config import (
|
from sglang.srt.layers.attention.qsa.config import (
|
||||||
QSA_VARIANT_COMPRESSED,
|
|
||||||
is_qwen_qsa,
|
is_qwen_qsa,
|
||||||
parse_qsa_profile,
|
parse_qsa_profile,
|
||||||
)
|
)
|
||||||
@@ -183,7 +182,6 @@ class QwenSparseAttnBackend(AttentionBackend):
|
|||||||
config = getattr(model_config, "hf_text_config", None)
|
config = getattr(model_config, "hf_text_config", None)
|
||||||
if config is None:
|
if config is None:
|
||||||
config = getattr(model_config, "hf_config", 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.qsa_profile = parse_qsa_profile(config)
|
||||||
self.max_context_len = int(getattr(model_config, "context_len", 0))
|
self.max_context_len = int(getattr(model_config, "context_len", 0))
|
||||||
self.compress_ratio = (
|
self.compress_ratio = (
|
||||||
@@ -631,20 +629,12 @@ class QwenSparseAttnBackend(AttentionBackend):
|
|||||||
f"mapping={token_to_batch_idx.numel()}, "
|
f"mapping={token_to_batch_idx.numel()}, "
|
||||||
f"positions={num_position_tokens}"
|
f"positions={num_position_tokens}"
|
||||||
)
|
)
|
||||||
write_locs = None
|
|
||||||
group_positions = None
|
|
||||||
group_sequence_ids = None
|
|
||||||
group_member_rows = None
|
|
||||||
decode_page_table = None
|
decode_page_table = None
|
||||||
decode_lengths = None
|
decode_lengths = None
|
||||||
decode_logical_positions = None
|
decode_logical_positions = None
|
||||||
pending_ring_slots = None
|
pending_ring_slots = None
|
||||||
compress_group_ring_locs = None
|
compress_group_ring_locs = None
|
||||||
extend_rope_matrix = 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 = (
|
write_locs, group_positions, group_sequence_ids, group_member_rows = (
|
||||||
self._qsa_build_write_plan(
|
self._qsa_build_write_plan(
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
@@ -835,17 +825,6 @@ class QwenSparseAttnBackend(AttentionBackend):
|
|||||||
]
|
]
|
||||||
self._extend_lens_pin_idx = 0
|
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(
|
def _capture_cuda_graph_metadata(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -856,7 +835,6 @@ class QwenSparseAttnBackend(AttentionBackend):
|
|||||||
forward_mode,
|
forward_mode,
|
||||||
spec_info,
|
spec_info,
|
||||||
) -> None:
|
) -> None:
|
||||||
self._require_compressed_cuda_graph_support()
|
|
||||||
self._require_chain_speculation(forward_mode, spec_info)
|
self._require_chain_speculation(forward_mode, spec_info)
|
||||||
if self.token_to_kv_pool is None:
|
if self.token_to_kv_pool is None:
|
||||||
self.token_to_kv_pool = getattr(self.runner, "token_to_kv_pool", None)
|
self.token_to_kv_pool = getattr(self.runner, "token_to_kv_pool", None)
|
||||||
|
|||||||
@@ -1887,25 +1887,16 @@ class KVCacheConfigurator:
|
|||||||
else mha_pool_class
|
else mha_pool_class
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.qsa.config import (
|
from sglang.srt.layers.attention.qsa.config import (
|
||||||
QSA_VARIANT_TOKENWISE,
|
|
||||||
parse_qsa_profile,
|
parse_qsa_profile,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.qsa_kv_pool import (
|
from sglang.srt.mem_cache.qsa_kv_pool import (
|
||||||
QSATokenToKVPool,
|
QSATokenToKVPool,
|
||||||
QwenDSATokenToKVPool,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
qsa_profile = parse_qsa_profile(self.model_config.hf_config)
|
qsa_profile = parse_qsa_profile(self.model_config.hf_config)
|
||||||
if qsa_profile is None:
|
if qsa_profile is None:
|
||||||
pool_class = HybridLinearKVPool
|
pool_class = HybridLinearKVPool
|
||||||
extra_args["use_mla"] = self.use_mla_backend
|
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:
|
else:
|
||||||
pool_class = QSATokenToKVPool
|
pool_class = QSATokenToKVPool
|
||||||
extra_args.update(
|
extra_args.update(
|
||||||
|
|||||||
@@ -283,107 +283,3 @@ class QSATokenToKVPool(HybridLinearKVPool):
|
|||||||
+ self.qsa_rope_position_buffer.numel() * 8
|
+ self.qsa_rope_position_buffer.numel() * 8
|
||||||
)
|
)
|
||||||
return k_size + qsa_k_size, v_size
|
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
|
|
||||||
|
|||||||
@@ -431,12 +431,10 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
@staticmethod
|
@staticmethod
|
||||||
def _compute_qsa_cell_size(*, hf_config, num_layers: int) -> int:
|
def _compute_qsa_cell_size(*, hf_config, num_layers: int) -> int:
|
||||||
from sglang.srt.layers.attention.qsa.config import (
|
from sglang.srt.layers.attention.qsa.config import (
|
||||||
QSA_VARIANT_COMPRESSED,
|
|
||||||
parse_qsa_profile,
|
parse_qsa_profile,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.qsa_kv_pool import (
|
from sglang.srt.mem_cache.qsa_kv_pool import (
|
||||||
QSATokenToKVPool,
|
QSATokenToKVPool,
|
||||||
QwenDSATokenToKVPool,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if num_layers == 0:
|
if num_layers == 0:
|
||||||
@@ -444,18 +442,12 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
qsa_profile = parse_qsa_profile(hf_config)
|
qsa_profile = parse_qsa_profile(hf_config)
|
||||||
if qsa_profile is None:
|
if qsa_profile is None:
|
||||||
return 0
|
return 0
|
||||||
if qsa_profile.variant == QSA_VARIANT_COMPRESSED:
|
|
||||||
return QSATokenToKVPool.qsa_bytes_per_token(
|
return QSATokenToKVPool.qsa_bytes_per_token(
|
||||||
kv_heads=qsa_profile.kv_heads,
|
kv_heads=qsa_profile.kv_heads,
|
||||||
head_dim=qsa_profile.head_dim,
|
head_dim=qsa_profile.head_dim,
|
||||||
compress_ratio=qsa_profile.compress_ratio,
|
compress_ratio=qsa_profile.compress_ratio,
|
||||||
num_layers=num_layers,
|
num_layers=num_layers,
|
||||||
)
|
)
|
||||||
return QwenDSATokenToKVPool.qsa_bytes_per_token(
|
|
||||||
kv_heads=qsa_profile.kv_heads,
|
|
||||||
head_dim=qsa_profile.head_dim,
|
|
||||||
num_layers=num_layers,
|
|
||||||
)
|
|
||||||
|
|
||||||
def _compute_dsa_indexer_cell_size(
|
def _compute_dsa_indexer_cell_size(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
from typing import Optional
|
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.runtime_context import attention_backends, get_spec
|
||||||
from sglang.srt.utils.common import (
|
from sglang.srt.utils.common import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
@@ -175,10 +175,6 @@ class DraftBackendFactory:
|
|||||||
backend.decode_attention_backend_str = "qsa"
|
backend.decode_attention_backend_str = "qsa"
|
||||||
|
|
||||||
def _create_qwen_qsa_draft_extend_backend(self):
|
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 (
|
from sglang.srt.layers.attention.qwen_sparse_attn_backend import (
|
||||||
QwenSparseAttnBackend,
|
QwenSparseAttnBackend,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ import torch
|
|||||||
from sglang.kernels.ops.attention import qwen38_qsa_sm121_varlen
|
from sglang.kernels.ops.attention import qwen38_qsa_sm121_varlen
|
||||||
from sglang.srt.configs.qwen4_exp import Qwen4ExpConfig
|
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 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 import qsa_indexer as qsa_indexer_module
|
||||||
from sglang.srt.layers.attention.qsa.kernel import (
|
from sglang.srt.layers.attention.qsa.kernel import (
|
||||||
expand_qsa_block_indices,
|
expand_qsa_block_indices,
|
||||||
@@ -261,22 +260,8 @@ def _compressed_config_namespace(**overrides):
|
|||||||
return SimpleNamespace(**fields)
|
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():
|
def test_qsa_profile_parses_compressed_qwen4_exp_schema():
|
||||||
from sglang.srt.layers.attention.qsa.config import (
|
from sglang.srt.layers.attention.qsa.config import (
|
||||||
QSA_ROPE_MROPE,
|
|
||||||
QSA_VARIANT_COMPRESSED,
|
|
||||||
is_qwen_qsa,
|
is_qwen_qsa,
|
||||||
parse_qsa_profile,
|
parse_qsa_profile,
|
||||||
)
|
)
|
||||||
@@ -297,14 +282,12 @@ def test_qsa_profile_parses_compressed_qwen4_exp_schema():
|
|||||||
)
|
)
|
||||||
for config in (wrapped, _compressed_config_namespace()):
|
for config in (wrapped, _compressed_config_namespace()):
|
||||||
profile = parse_qsa_profile(config)
|
profile = parse_qsa_profile(config)
|
||||||
assert profile.variant == QSA_VARIANT_COMPRESSED
|
|
||||||
assert profile.n_heads == 8
|
assert profile.n_heads == 8
|
||||||
assert profile.kv_heads == 1
|
assert profile.kv_heads == 1
|
||||||
assert profile.head_dim == 128
|
assert profile.head_dim == 128
|
||||||
assert profile.budget == TOKEN_TOPK
|
assert profile.budget == TOKEN_TOPK
|
||||||
assert profile.compress_ratio == COMPRESS_RATIO
|
assert profile.compress_ratio == COMPRESS_RATIO
|
||||||
assert profile.block_topk == BLOCK_TOPK
|
assert profile.block_topk == BLOCK_TOPK
|
||||||
assert profile.rope_mode == QSA_ROPE_MROPE
|
|
||||||
assert is_qwen_qsa(config)
|
assert is_qwen_qsa(config)
|
||||||
# The legacy backend module keeps re-exporting the shared detector.
|
# The legacy backend module keeps re-exporting the shared detector.
|
||||||
assert qsa_backend_module.is_qwen_qsa is is_qwen_qsa
|
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")
|
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
|
from sglang.srt.layers.attention.qsa.glue import build_qsa_indexer
|
||||||
|
|
||||||
recorded = {}
|
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
|
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:
|
try:
|
||||||
build_qsa_indexer(SimpleNamespace(), layer_id=0, rotary_emb=rotary)
|
build_qsa_indexer(SimpleNamespace(), layer_id=0, rotary_emb=rotary)
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
@@ -438,9 +403,6 @@ def test_qsa_draft_extend_backend_decision_follows_profile():
|
|||||||
assert isinstance(backend, QwenSparseAttnBackend)
|
assert isinstance(backend, QwenSparseAttnBackend)
|
||||||
assert backend.runner is compressed.draft_model_runner
|
assert backend.runner is compressed.draft_model_runner
|
||||||
assert backend.decode_attention_backend_str == "qsa"
|
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):
|
def _make_mtp_draft_batch(steps: int, seq_lens=(8, 16), loc_base: int = 40):
|
||||||
|
|||||||
Reference in New Issue
Block a user