{"record":{"id":"3c42450c6aa2b838","repo":"jax-ml/jax","slug":"pages-per-seq-must-be-greater-or-equal-to-min","errorCode":null,"errorMessage":"{pages_per_seq=} must be greater or equal to {min_pages_per_seq=} given {max_kv_len=} and {page_size=}.","messagePattern":"(.+?) must be greater or equal to (.+?) given (.+?) and (.+?)\\.","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/experimental/pallas/ops/tpu/ragged_paged_attention/kernel.py","lineNumber":197,"sourceCode":"      sm_scale=sm_scale,\n      sliding_window=sliding_window,\n      soft_cap=soft_cap,\n      mask_value=mask_value,\n      k_scale=k_scale,\n      v_scale=v_scale,\n      num_kv_pages_per_block=num_kv_pages_per_block,\n      num_queries_per_block=num_queries_per_block,\n      vmem_limit_bytes=vmem_limit_bytes,\n  )\n  max_num_batched_tokens = q.shape[0]\n  page_size = kv_pages.shape[1]\n  max_num_seqs, pages_per_seq = page_indices.shape\n  if num_seqs[0] > max_num_seqs:\n    raise ValueError(f\"{num_seqs[0]=} must be less or equal to {max_num_seqs=}\")\n  max_kv_len = jnp.max(kv_lens)\n  min_pages_per_seq = pl.cdiv(max_kv_len, page_size)\n  if pages_per_seq < min_pages_per_seq:\n    raise ValueError(\n        f\"{pages_per_seq=} must be greater or equal to\"\n        f\" {min_pages_per_seq=} given {max_kv_len=} and {page_size=}.\"\n    )\n  if cu_q_lens[num_seqs[0]] > max_num_batched_tokens:\n    raise ValueError(\n        f\"Total q tokens {cu_q_lens[num_seqs[0]]} must be less or equal to\"\n        f\" {max_num_batched_tokens=}.\"\n    )\n  for i in range(num_seqs[0]):\n    q_len = cu_q_lens[i + 1] - cu_q_lens[i]\n    kv_len = kv_lens[i]\n    if q_len > kv_len:\n      raise ValueError(\n          f\"{q_len=} must be less or equal to {kv_len=} at sequence {i}.\"\n      )\n\n\n# Expect to run these checks during compile time.","sourceCodeStart":179,"sourceCodeEnd":215,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/pallas/ops/tpu/ragged_paged_attention/kernel.py#L179-L215","documentation":"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.","triggerScenarios":"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.","commonSituations":"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.","solutions":["Grow pages_per_seq to at least ceil(max(kv_lens)/page_size)","Trim/quantize sequences so max kv_len fits the existing pages_per_seq","Compute and assert pages_per_seq >= -(-max_kv_len // page_size) before the call"],"exampleFix":"// before\npage_indices = jnp.zeros((max_seqs, 16), jnp.int32)  # kv_len up to 4096, page 128\n// after\npage_indices = jnp.zeros((max_seqs, 32), jnp.int32)  # ceil(4096/128)","handlingStrategy":"validation","validationCode":"min_pages = -(-int(jnp.max(kv_lens)) // kv_pages.shape[1])  # ceil\nassert page_indices.shape[1] >= min_pages","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Allocate pages_per_seq from max context length: ceil(max_len/page_size)","Re-check capacity whenever max sequence length or page size changes"],"tags":["jax","pallas","tpu","ragged-attention","kv-cache","capacity-validation"],"backgroundTag":"capacity-exceeded","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}