jax-ml/jax · error · ValueError
{num_seqs[0]=} must be less or equal to {max_num_seqs=}
Error message
{num_seqs[0]=} must be less or equal to {max_num_seqs=} What it means
In the ragged paged-attention kernel, num_seqs is a device array holding the live number of sequences, while page_indices is preallocated to max_num_seqs rows. The dynamic (compile-time-run) validation enforces num_seqs[0] <= max_num_seqs so the kernel never reads past the page table.
Source
Thrown at jax/experimental/pallas/ops/tpu/ragged_paged_attention/kernel.py:193
kv_lens,
page_indices,
cu_q_lens,
num_seqs,
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}."View on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Cap admitted sequences: num_seqs = min(num_seqs, page_indices.shape[0])
- Preallocate page_indices with max_num_seqs >= scheduler max batch size
- Validate num_seqs[0] <= page_indices.shape[0] before launching
Example fix
// before page_indices = jnp.zeros((32, pages), jnp.int32) ragged_paged_attention(q, k_pages, page_indices, num_seqs=jnp.array([40]), ...) // after page_indices = jnp.zeros((64, pages), jnp.int32) # >= max batch
Defensive patterns
Strategy: validation
Validate before calling
max_num_seqs = page_indices.shape[0] assert int(num_seqs[0]) <= max_num_seqs, (num_seqs[0], max_num_seqs)
Prevention
- Size page_indices rows to the scheduler's max batch, not current batch
- Centralize capacity checks in the scheduler admission step
When it happens
Trigger: Calling ragged_paged_attention (or its _test helper via dynamic_validate_inputs) with num_seqs[0] greater than page_indices.shape[0], e.g. scheduling 40 sequences into a page table sized for 32.
Common situations: Continuous batching schedulers that admit more sequences than the preallocated KV-cache capacity; miscomputing max_num_seqs when allocating page_indices; shrinking cache size while keeping scheduler limits unchanged.
Related errors
- {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
- {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/a38d1afc609fb38a.
Report an issue: GitHub.