{"record":{"id":"ba0049a2f0d029a7","repo":"jax-ml/jax","slug":"num-kv-pages-per-block-must-be-in-range-0-pa","errorCode":null,"errorMessage":"{num_kv_pages_per_block=} must be in range (0, {pages_per_seq}].","messagePattern":"(.+?) must be in range \\(0, (.+?)\\]\\.","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/experimental/pallas/ops/tpu/ragged_paged_attention/kernel.py","lineNumber":279,"sourceCode":"      or page_indices.dtype != jnp.int32\n      or cu_q_lens.dtype != jnp.int32\n  ):\n    raise ValueError(\n        \"The dtype of `kv_lens`, `page_indices`, and `cu_q_lens` must be\"\n        f\" int32. Got {kv_lens.dtype=}, {page_indices.dtype=},\"\n        f\" {cu_q_lens.dtype=}.\"\n    )\n  if num_q_heads % num_kv_heads != 0:\n    raise ValueError(f\"{num_q_heads=} must be divisible by {num_kv_heads=}\")\n  if sliding_window is not None and sliding_window <= 0:\n    raise ValueError(f\"{sliding_window=} must be positive.\")\n  if soft_cap is not None and soft_cap == 0.0:\n    raise ValueError(f\"{soft_cap=} must not be 0.0.\")\n  if (\n      num_kv_pages_per_block is not None\n      and not 0 < num_kv_pages_per_block <= pages_per_seq\n  ):\n    raise ValueError(\n        f\"{num_kv_pages_per_block=} must be in range (0, {pages_per_seq}].\"\n    )\n  if num_queries_per_block is not None and num_queries_per_block <= 0:\n    raise ValueError(f\"{num_queries_per_block=} must be positive.\")\n  if vmem_limit_bytes is not None and vmem_limit_bytes <= 0:\n    raise ValueError(f\"{vmem_limit_bytes=} must be positive.\")\n  del sm_scale  # No constraints on sm_scale.\n  del mask_value  # No consstraints on mask_value.\n\n\ndef ragged_paged_attention_kernel(\n    # Prefetch\n    kv_lens_ref,  # [max_num_seqs]\n    page_indices_ref,  # [max_num_seqs, pages_per_seq]\n    cu_q_lens_ref,  # [max_num_seqs + 1]\n    seq_buf_idx_ref,\n    # TODO(jevinjiang): if OOM in SMEM, consider pack to other scalar refs.\n    num_seqs_ref,","sourceCodeStart":261,"sourceCodeEnd":297,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/pallas/ops/tpu/ragged_paged_attention/kernel.py#L261-L297","documentation":"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.","triggerScenarios":"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.","commonSituations":"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.","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"],"exampleFix":"// before\nattn(..., num_kv_pages_per_block=128)  # pages_per_seq is 64\n// after\nnum_kv_pages_per_block = min(128, pages_per_seq)  # or None\nattn(..., num_kv_pages_per_block=num_kv_pages_per_block)","handlingStrategy":"validation","validationCode":"pages_per_seq = page_indices.shape[1]\nif num_kv_pages_per_block is not None:\n    assert 0 < num_kv_pages_per_block <= pages_per_seq","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Derive pages_per_seq from the actual page table, not a stale constant","Prefer None (autotuned) unless you have measured a better block size"],"tags":["jax","pallas","tpu","block-size","paged-attention"],"backgroundTag":"invalid-argument-value","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}