{"record":{"id":"4848dd6ef10b514a","repo":"jax-ml/jax","slug":"num-seqs-shape-must-be-1","errorCode":null,"errorMessage":"{num_seqs.shape=} must be (1,)","messagePattern":"(.+?) must be \\(1,\\)","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/experimental/pallas/ops/tpu/ragged_paged_attention/kernel.py","lineNumber":244,"sourceCode":"    sliding_window: int | None = None,\n    soft_cap: float | None = None,\n    mask_value: float | None = None,\n    k_scale: float | None = None,\n    v_scale: float | None = None,\n    # Kernel tuning params.\n    num_kv_pages_per_block: int | None = None,\n    num_queries_per_block: int | None = None,\n    vmem_limit_bytes: int | None = None,\n):\n  _, num_q_heads, head_dim = q.shape\n  _, _, num_combined_kv_heads, head_dim_k = kv_pages.shape\n  assert num_combined_kv_heads % 2 == 0\n  assert isinstance(k_scale, float) or k_scale is None\n  assert isinstance(v_scale, float) or v_scale is None\n  num_kv_heads = num_combined_kv_heads // 2\n  max_num_seqs, pages_per_seq = page_indices.shape\n  if num_seqs.shape != (1,):\n    raise ValueError(f\"{num_seqs.shape=} must be (1,)\")\n  if head_dim_k != head_dim:\n    raise ValueError(\n        f\"Q head_dim {head_dim} must be the same as that of K/V {head_dim_k}.\"\n    )\n  if kv_lens.shape != (max_num_seqs,):\n    raise ValueError(\n        f\"Expected {kv_lens.shape=} to be ({max_num_seqs},) where\"\n        \" `max_num_seqs` is `page_indices.shape[0]`.\"\n    )\n  if cu_q_lens.shape != (max_num_seqs + 1,):\n    raise ValueError(\n        f\"Expected {cu_q_lens.shape=} to be ({max_num_seqs + 1},)  where\"\n        \" `max_num_seqs` is `page_indices.shape[0]`.\"\n    )\n  if (\n      kv_lens.dtype != jnp.int32\n      or page_indices.dtype != jnp.int32\n      or cu_q_lens.dtype != jnp.int32","sourceCodeStart":226,"sourceCodeEnd":262,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/pallas/ops/tpu/ragged_paged_attention/kernel.py#L226-L262","documentation":"static_validate_inputs runs at trace/compile time on the ragged paged-attention path (used by ragged_paged_attention, ref impl, and dynamic validation). num_seqs must be a shape-(1,) device array holding the live sequence count; other shapes (scalar, (n,), etc.) are rejected because the kernel reads num_seqs[0] as a single value.","triggerScenarios":"Passing num_seqs as a Python int, a 0-d array, or a per-sequence array of counts instead of jnp.array([n], dtype=int). Note this must be an array even though the count is dynamic.","commonSituations":"Passing num_seqs=5 (int) for convenience; reshaping scheduler state; migrating from an API that accepted a scalar count.","solutions":["Wrap the count: num_seqs = jnp.array([n], dtype=jnp.int32)","Keep scheduler output shapes stable: num_seqs always shape (1,)"],"exampleFix":"// before\nragged_paged_attention(q, k_pages, page_indices, num_seqs=5, ...)\n// after\nragged_paged_attention(q, k_pages, page_indices, num_seqs=jnp.array([5], jnp.int32), ...)","handlingStrategy":"type-guard","validationCode":"num_seqs = jnp.asarray(num_seqs).reshape(1)","typeGuard":"def is_num_seqs_shape(a): return isinstance(a, jax.Array) and a.shape == (1,)","tryCatchPattern":null,"preventionTips":["Always construct num_seqs as jnp.array([n], jnp.int32)","Keep scheduler state shapes fixed across steps"],"tags":["jax","pallas","tpu","ragged-attention","shape-validation"],"backgroundTag":"shape-validation-failed","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}