jax-ml/jax · error · ValueError
{pages_per_seq=} must be greater or equal to {min_pages_per_
Error message
{pages_per_seq=} must be greater or equal to {min_pages_per_seq=} given {max_kv_len=} and {page_size=}. What it means
Each sequence needs ceil(max_kv_len / page_size) pages to hold its KV cache. The dynamic validation computes min_pages_per_seq from the largest kv_lens entry and requires page_indices.shape[1] (pages per sequence) to be at least that, otherwise pages would be missing and attention results silently truncated, so it raises instead.
Source
Thrown at jax/experimental/pallas/ops/tpu/ragged_paged_attention/kernel.py:197
sm_scale=sm_scale,
sliding_window=sliding_window,
soft_cap=soft_cap,
mask_value=mask_value,
k_scale=k_scale,
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.View on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Grow pages_per_seq to at least ceil(max(kv_lens)/page_size)
- Trim/quantize sequences so max kv_len fits the existing pages_per_seq
- Compute and assert pages_per_seq >= -(-max_kv_len // page_size) before the call
Example fix
// before page_indices = jnp.zeros((max_seqs, 16), jnp.int32) # kv_len up to 4096, page 128 // after page_indices = jnp.zeros((max_seqs, 32), jnp.int32) # ceil(4096/128)
Defensive patterns
Strategy: validation
Validate before calling
min_pages = -(-int(jnp.max(kv_lens)) // kv_pages.shape[1]) # ceil assert page_indices.shape[1] >= min_pages
Prevention
- Allocate pages_per_seq from max context length: ceil(max_len/page_size)
- Re-check capacity whenever max sequence length or page size changes
When it happens
Trigger: Calling ragged_paged_attention where page_indices has too few page slots for the longest sequence, e.g. pages_per_seq=16, page_size=128 but one sequence has kv_len=4096 needing 32 pages.
Common situations: Long-context prompts exceeding the allocated KV budget; increasing max sequence length without growing the paged cache; shrinking page count to save memory while retaining long prompts.
Related errors
- {num_seqs[0]=} must be less or equal to {max_num_seqs=}
- Total q tokens {cu_q_lens[num_seqs[0]]} must be less or equa
- {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/3c42450c6aa2b838.
Report an issue: GitHub.