jax-ml/jax · error · ValueError
{num_kv_pages_per_block=} must be in range (0, {pages_per_se
Error message
{num_kv_pages_per_block=} must be in range (0, {pages_per_seq}]. What it means
num_kv_pages_per_block optionally overrides the KV block size of the ragged paged attention kernel, but must satisfy 0 < num_kv_pages_per_block <= pages_per_seq (the total number of pages per sequence in the page table). Values outside that range cannot be tiled and are rejected by static_validate_inputs.
Source
Thrown at jax/experimental/pallas/ops/tpu/ragged_paged_attention/kernel.py:279
or page_indices.dtype != jnp.int32
or cu_q_lens.dtype != jnp.int32
):
raise ValueError(
"The dtype of `kv_lens`, `page_indices`, and `cu_q_lens` must be"
f" int32. Got {kv_lens.dtype=}, {page_indices.dtype=},"
f" {cu_q_lens.dtype=}."
)
if num_q_heads % num_kv_heads != 0:
raise ValueError(f"{num_q_heads=} must be divisible by {num_kv_heads=}")
if sliding_window is not None and sliding_window <= 0:
raise ValueError(f"{sliding_window=} must be positive.")
if soft_cap is not None and soft_cap == 0.0:
raise ValueError(f"{soft_cap=} must not be 0.0.")
if (
num_kv_pages_per_block is not None
and not 0 < num_kv_pages_per_block <= pages_per_seq
):
raise ValueError(
f"{num_kv_pages_per_block=} must be in range (0, {pages_per_seq}]."
)
if num_queries_per_block is not None and num_queries_per_block <= 0:
raise ValueError(f"{num_queries_per_block=} must be positive.")
if vmem_limit_bytes is not None and vmem_limit_bytes <= 0:
raise ValueError(f"{vmem_limit_bytes=} must be positive.")
del sm_scale # No constraints on sm_scale.
del mask_value # No consstraints on mask_value.
def ragged_paged_attention_kernel(
# Prefetch
kv_lens_ref, # [max_num_seqs]
page_indices_ref, # [max_num_seqs, pages_per_seq]
cu_q_lens_ref, # [max_num_seqs + 1]
seq_buf_idx_ref,
# TODO(jevinjiang): if OOM in SMEM, consider pack to other scalar refs.
num_seqs_ref,View on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Pass None to let the kernel/autotuner pick the block size
- Otherwise clamp to at most pages_per_seq (number of page-table columns per sequence)
- Recompute pages_per_seq = ceil(max_kv_len / page_size) for your paging config before validating
Example fix
// before attn(..., num_kv_pages_per_block=128) # pages_per_seq is 64 // after num_kv_pages_per_block = min(128, pages_per_seq) # or None attn(..., num_kv_pages_per_block=num_kv_pages_per_block)
Defensive patterns
Strategy: validation
Validate before calling
pages_per_seq = page_indices.shape[1]
if num_kv_pages_per_block is not None:
assert 0 < num_kv_pages_per_block <= pages_per_seq Prevention
- Derive pages_per_seq from the actual page table, not a stale constant
- Prefer None (autotuned) unless you have measured a better block size
When it happens
Trigger: Passing a KV block size larger than pages_per_seq (e.g. 64 when the page table only has 32 pages per sequence), or <= 0.
Common situations: Reusing autotuned or hand-tuned block sizes from a different model/page-size configuration (e.g. page_size 16 vs 128 changes pages_per_seq) after changing the paging setup.
Understand the failure class
Background: "Must be a positive integer", "Invalid value", "Unsupported": the invalid-argument-value error family, when a library rejects the value you pass — this error's family across 35 libraries.
Related errors
- k_pages and v_pages must have the same shape. Got {k_pages.s
- Number of Q heads must be divisible by number of KV heads. G
- head_dim of Q must be the same as that of K/V. Got {head_dim
- pages_per_compute_block must be divisible by pages per seque
- `lengths` and `q` must have the same batch size
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/ba0049a2f0d029a7.
Report an issue: GitHub.