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

  1. Grow pages_per_seq to at least ceil(max(kv_lens)/page_size)
  2. Trim/quantize sequences so max kv_len fits the existing pages_per_seq
  3. 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

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


AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27). Data as JSON: /api/errors/3c42450c6aa2b838. Report an issue: GitHub.