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

  1. Pass None to let the kernel/autotuner pick the block size
  2. Otherwise clamp to at most pages_per_seq (number of page-table columns per sequence)
  3. 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

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


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