jax-ml/jax · error · ValueError
{q_len=} must be less or equal to {kv_len=} at sequence {i}.
Error message
{q_len=} must be less or equal to {kv_len=} at sequence {i}. What it means
The ragged paged-attention kernel requires each sequence's query length to not exceed its KV length (standard causal prefill/decode constraint used for masking). The dynamic validator loops over sequences comparing cu_q_lens[i+1]-cu_q_lens[i] against kv_lens[i] and raises naming the offending sequence index.
Source
Thrown at jax/experimental/pallas/ops/tpu/ragged_paged_attention/kernel.py:210
if num_seqs[0] > max_num_seqs:
raise ValueError(f"{num_seqs[0]=} must be less or equal to {max_num_seqs=}")
max_kv_len = jnp.max(kv_lens)
min_pages_per_seq = pl.cdiv(max_kv_len, page_size)
if pages_per_seq < min_pages_per_seq:
raise ValueError(
f"{pages_per_seq=} must be greater or equal to"
f" {min_pages_per_seq=} given {max_kv_len=} and {page_size=}."
)
if cu_q_lens[num_seqs[0]] > max_num_batched_tokens:
raise ValueError(
f"Total q tokens {cu_q_lens[num_seqs[0]]} must be less or equal to"
f" {max_num_batched_tokens=}."
)
for i in range(num_seqs[0]):
q_len = cu_q_lens[i + 1] - cu_q_lens[i]
kv_len = kv_lens[i]
if q_len > kv_len:
raise ValueError(
f"{q_len=} must be less or equal to {kv_len=} at sequence {i}."
)
# Expect to run these checks during compile time.
def static_validate_inputs(
q: jax.Array, # [max_num_batched_tokens, num_q_heads, head_dim]
kv_pages: jax.Array, # [total_num_pages, page_size, num_combined_kv_heads, head_dim]
kv_lens: jax.Array, # i32[max_num_seqs]
page_indices: jax.Array, # i32[max_num_seqs, pages_per_seq]
cu_q_lens: jax.Array, # i32[max_num_seqs + 1]
num_seqs: jax.Array, # i32[1]
*,
# These inputs are optional. If not specified, we will not validate them.
sm_scale: float | None = None,
sliding_window: int | None = None,
soft_cap: float | None = None,
mask_value: float | None = None,View on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Fix kv_lens[i] to include the new query tokens (append them to the cache first)
- Chunk queries so each chunk's q_len <= kv_len_i
- Validate per-sequence q_len <= kv_len in the scheduler before launch
Example fix
// before kv_lens = old_kv_lens # stale, excludes new tokens // after kv_lens = old_kv_lens + current_q_lens # cache updated with new tokens
Defensive patterns
Strategy: validation
Validate before calling
for i in range(int(num_seqs[0])):
q_len = int(cu_q_lens[i+1] - cu_q_lens[i])
assert q_len <= int(kv_lens[i]), (i, q_len, int(kv_lens[i])) Prevention
- Update kv_lens whenever tokens are appended to the cache
- In chunked prefill, chunk size must not exceed cached prefix + chunk
When it happens
Trigger: Calling ragged_paged_attention with a sequence that has more query tokens than cached KV tokens, e.g. q_len=512 but kv_len=256 for sequence i; commonly from wrong kv_lens after cache eviction or chunked prefill bookkeeping.
Common situations: Chunked prefill where query chunk exceeds already-cached KV for the chunk's prefix; kv_lens not updated after appending new tokens; off-by-one in cu_q_lens construction.
Related errors
- {num_seqs[0]=} must be less or equal to {max_num_seqs=}
- {pages_per_seq=} must be greater or equal to {min_pages_per_
- Total q tokens {cu_q_lens[num_seqs[0]]} must be less or equa
- {num_seqs.shape=} must be (1,)
- Q head_dim {head_dim} must be the same as that of K/V {head_
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/453d20457508aaed.
Report an issue: GitHub.