jax-ml/jax · error · ValueError
Total q tokens {cu_q_lens[num_seqs[0]]} must be less or equa
Error message
Total q tokens {cu_q_lens[num_seqs[0]]} must be less or equal to {max_num_batched_tokens=}. What it means
cu_q_lens is the cumulative query-token count (like cu_seqlens in flash-attention) and its final relevant entry cu_q_lens[num_seqs[0]] gives total queries. q is padded to max_num_batched_tokens rows, so the total must fit within q.shape[0]; otherwise the kernel would gather query rows beyond the buffer.
Source
Thrown at jax/experimental/pallas/ops/tpu/ragged_paged_attention/kernel.py:202
v_scale=v_scale,
num_kv_pages_per_block=num_kv_pages_per_block,
num_queries_per_block=num_queries_per_block,
vmem_limit_bytes=vmem_limit_bytes,
)
max_num_batched_tokens = q.shape[0]
page_size = kv_pages.shape[1]
max_num_seqs, pages_per_seq = page_indices.shape
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]View on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Allocate q with q.shape[0] >= int(cu_q_lens[num_seqs]) and pad remaining rows with zeros
- Rebuild cu_q_lens from the actual packed batch: cu = jnp.concatenate([jnp.array([0]), jnp.cumsum(q_lens)])
- Validate cu_q_lens[-1 relevant] <= q.shape[0] before launch
Example fix
// before q = q_tokens[:1024] # but cu_q_lens[num_seqs] == 1200 // after q = jnp.pad(q_tokens[:1200], ((0, 88), (0,0), (0,0))) # fit budget
Defensive patterns
Strategy: validation
Validate before calling
total_q = int(cu_q_lens[int(num_seqs[0])]) assert total_q <= q.shape[0], (total_q, q.shape[0])
Prevention
- Build cu_q_lens with jnp.cumsum from the actual packed lengths
- Pad q to the token budget and mask with lengths=0 rows
When it happens
Trigger: Calling ragged_paged_attention where q was sliced to fewer rows than sum of per-sequence query lengths, or cu_q_lens was built against a different token budget than the actual q allocation.
Common situations: Chunked prefill schedulers that batch tokens up to a budget but build cu_q_lens against a larger one; off-by-one in cumulative sums; padding q to the wrong axis length.
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_
- {q_len=} must be less or equal to {kv_len=} at sequence {i}.
- {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/0cb44620c8775179.
Report an issue: GitHub.