{"record":{"id":"453d20457508aaed","repo":"jax-ml/jax","slug":"q-len-must-be-less-or-equal-to-kv-len-at-seq","errorCode":null,"errorMessage":"{q_len=} must be less or equal to {kv_len=} at sequence {i}.","messagePattern":"(.+?) must be less or equal to (.+?) at sequence (.+?)\\.","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/experimental/pallas/ops/tpu/ragged_paged_attention/kernel.py","lineNumber":210,"sourceCode":"  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.\ndef static_validate_inputs(\n    q: jax.Array,  # [max_num_batched_tokens, num_q_heads, head_dim]\n    kv_pages: jax.Array,  # [total_num_pages, page_size, num_combined_kv_heads, head_dim]\n    kv_lens: jax.Array,  # i32[max_num_seqs]\n    page_indices: jax.Array,  # i32[max_num_seqs, pages_per_seq]\n    cu_q_lens: jax.Array,  # i32[max_num_seqs + 1]\n    num_seqs: jax.Array,  # i32[1]\n    *,\n    # These inputs are optional. If not specified, we will not validate them.\n    sm_scale: float | None = None,\n    sliding_window: int | None = None,\n    soft_cap: float | None = None,\n    mask_value: float | None = None,","sourceCodeStart":192,"sourceCodeEnd":228,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/pallas/ops/tpu/ragged_paged_attention/kernel.py#L192-L228","documentation":"The ragged paged-attention kernel requires each sequence's query length to not exceed its KV length (standard causal prefill/decode constraint used for masking). The dynamic validator loops over sequences comparing cu_q_lens[i+1]-cu_q_lens[i] against kv_lens[i] and raises naming the offending sequence index.","triggerScenarios":"Calling ragged_paged_attention with a sequence that has more query tokens than cached KV tokens, e.g. q_len=512 but kv_len=256 for sequence i; commonly from wrong kv_lens after cache eviction or chunked prefill bookkeeping.","commonSituations":"Chunked prefill where query chunk exceeds already-cached KV for the chunk's prefix; kv_lens not updated after appending new tokens; off-by-one in cu_q_lens construction.","solutions":["Fix kv_lens[i] to include the new query tokens (append them to the cache first)","Chunk queries so each chunk's q_len <= kv_len_i","Validate per-sequence q_len <= kv_len in the scheduler before launch"],"exampleFix":"// before\nkv_lens = old_kv_lens                       # stale, excludes new tokens\n// after\nkv_lens = old_kv_lens + current_q_lens      # cache updated with new tokens","handlingStrategy":"validation","validationCode":"for i in range(int(num_seqs[0])):\n    q_len = int(cu_q_lens[i+1] - cu_q_lens[i])\n    assert q_len <= int(kv_lens[i]), (i, q_len, int(kv_lens[i]))","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Update kv_lens whenever tokens are appended to the cache","In chunked prefill, chunk size must not exceed cached prefix + chunk"],"tags":["jax","pallas","tpu","ragged-attention","length-validation"],"backgroundTag":"invalid-sequence-lengths","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}