sgl-project/sglang · error · ValueError

Unsupported query shape for Quest: {queries.shape}

Error message

Unsupported query shape for Quest: {queries.shape}

What it means

Quest's _retrieve_page_scores only accepts 1-D-projected 2-D or 3-D query tensors; anything else (0-D, 4-D, etc.) raises ValueError with the offending shape. The algorithm must reshape or use queries directly as (bs, q_heads, head_dim).

Source

Thrown at python/sglang/srt/mem_cache/sparsity/algorithms/quest_algorithm.py:146

        phys_pages_clamped = phys_pages.clamp(0, self.page_k_min[layer_id].shape[0] - 1)

        k_min = self.page_k_min[layer_id][phys_pages_clamped]
        k_max = self.page_k_max[layer_id][phys_pages_clamped]
        valid_mask = self.page_valid[layer_id][phys_pages_clamped]
        # Align query shape to KV heads.
        head_dim = k_min.shape[-1]
        if queries.dim() == 2:
            bs, hidden = queries.shape
            if hidden % head_dim != 0:
                raise ValueError(
                    f"Quest query hidden size {hidden} not divisible by head_dim {head_dim}"
                )
            q_heads = hidden // head_dim
            q = queries.view(bs, q_heads, head_dim)
        elif queries.dim() == 3:
            q = queries
        else:
            raise ValueError(f"Unsupported query shape for Quest: {queries.shape}")

        kv_heads = k_min.shape[-2]
        q_heads = q.shape[1]
        if q_heads != kv_heads:
            if q_heads % kv_heads != 0:
                raise ValueError(
                    f"Query heads {q_heads} not divisible by KV heads {kv_heads}"
                )
            group = q_heads // kv_heads
            # Average grouped query heads to align with KV heads (approximation for MQA/GQA).
            q = q.view(q.shape[0], kv_heads, group, head_dim).mean(dim=2)

        q = q.to(k_min.dtype).unsqueeze(1)  # [bs, 1, kv_heads, head_dim]

        criticality = torch.where(q >= 0, q * k_max, q * k_min).sum(dim=(2, 3))
        criticality = torch.where(
            valid_mask, criticality, torch.full_like(criticality, float("-inf"))
        )

View on GitHub (pinned to 0132848349)

Solutions

  1. Squeeze/reshape queries to (bs, hidden) or (bs, q_heads, head_dim) before calling retrieve
  2. Index the specific layer: queries = full_q[:, layer_id] before retrieval

Example fix

# before
scores = quest._retrieve_page_scores(queries=q_all_layers, ...)  # 4-D
# after
scores = quest._retrieve_page_scores(queries=q_all_layers[:, layer_id], ...)  # 3-D
Defensive patterns

Strategy: type-guard

Validate before calling

assert queries.dim() in (2, 3), f"bad query rank: {queries.shape}"

Type guard

def quest_rank_ok(queries) -> bool:
    return queries.dim() in (2, 3)

Prevention

When it happens

Trigger: Passing queries with dim() not in (2, 3) — e.g. a 4-D (bs, layers, heads, dim) tensor or a 1-D flattened vector — to Quest retrieval.

Common situations: Adapter code forwarding raw multi-dimensional attention tensors straight into the sparse algorithm instead of the per-layer per-request queries.

Related errors


AI-assisted analysis of sgl-project/sglang@0132848349 (2026-08-28). Data as JSON: /api/errors/a8155e3f08886f72. Report an issue: GitHub.