{"record":{"id":"3cf080a0fecc1e01","repo":"jax-ml/jax","slug":"need-to-specify-backward-blocks-3cf080","errorCode":null,"errorMessage":"Need to specify backward blocks.","messagePattern":"Need to specify backward blocks\\.","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py","lineNumber":2279,"sourceCode":"    residual_checkpoint_name: str | None,\n    mask_function: MaskFunctionType | None,\n    attn_logits_soft_cap: float | None,\n    interpret: bool,\n    res: SplashResidualsType,\n    do: jax.Array,\n) -> tuple[\n    mask_info_lib.MaskInfo | None,  # fwd_mask_info\n    mask_info_lib.MaskInfo | None,  # dq_mask_info\n    mask_info_lib.MaskInfo | None,  # dvk_mask_info\n    jax.Array,  # q\n    jax.Array,  # k\n    jax.Array,  # v\n    SegmentIds | None,  # segmend_ids\n    jax.Array | None,  # sinks\n]:\n  del save_residuals, residual_checkpoint_name\n  if not block_sizes.has_backward_blocks:\n    raise ValueError(\"Need to specify backward blocks.\")\n  bq_dq, bkv_dq = block_sizes.block_q_dq, block_sizes.block_kv_dq\n  bq_dkv, bkv_dkv_memory, bkv_dkv_compute = (\n      block_sizes.block_q_dkv,\n      block_sizes.block_kv_dkv,\n      block_sizes.block_kv_dkv_compute,\n  )\n  use_fused_bwd_kernel = block_sizes.use_fused_bwd_kernel\n  (\n      q,\n      k,\n      v,\n      segment_ids,\n      sinks,\n      o,\n      logsumexp,\n      dq_mask_info,\n      dkv_mask_info,\n  ) = res","sourceCodeStart":2261,"sourceCodeEnd":2297,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py#L2261-L2297","documentation":"The splash attention backward primitive requires explicit backward block sizes; unlike some kernels they are not inferred. Calling jax.grad (or .with_backward()) on a splash kernel whose BlockSizes lacks backward blocks raises this immediately.","triggerScenarios":"Calling make_splash_attention(...) without with_backward(...) (or without passing backward BlockSizes) and then differentiating the function, which routes into _splash_attention_bwd.","commonSituations":"Reusing a forward-only kernel config for training; upgrading jax versions where defaults for backward blocks changed from implicit to explicit; constructing BlockSizes manually and leaving bwd fields None.","solutions":["Create the kernel via BlockSizes.fwd_bwd(...) or call .with_backward(...) on the forward kernel so backward blocks are populated","Pass explicit block_q_dq/block_kv_dq/block_q_dkv/block_kv_dkv when constructing BlockSizes for training","If you only need inference, avoid jax.grad over this function"],"exampleFix":"// before\nkernel = make_splash_attention(mask, block_sizes=BlockSizes(...fwd only...))\nloss = jax.grad(fn)(...)  # raises\n// after\nkernel = make_splash_attention(mask, block_sizes=BlockSizes.fwd_bwd(...))\nloss = jax.grad(fn)(...)","handlingStrategy":"validation","validationCode":"assert block_sizes.has_backward_blocks, 'need BlockSizes.fwd_bwd(...) for grad'","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Use .with_backward() whenever training","Build kernels via BlockSizes.fwd_bwd for training paths"],"tags":["jax","pallas","splash-attention","block-sizes","backward"],"backgroundTag":"missing-backward-config","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}