{"record":{"id":"779849394d602d30","repo":"jax-ml/jax","slug":"only-1-reduction-dimension-is-supported","errorCode":null,"errorMessage":"Only 1 reduction dimension is supported.","messagePattern":"Only 1 reduction dimension is supported\\.","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/experimental/mosaic/gpu/dialect_lowering.py","lineNumber":889,"sourceCode":"      raise NotImplementedError(f\"Unsupported reduction kind: {op.kind}\")\n  assert isinstance(result.layout, fa.WGSplatFragLayout)\n  return [result.registers.item()]\n\n\n@_register_lowering(vector.MultiDimReductionOp)\ndef _vector_multi_dim_reduction_op_lowering_rule(\n    ctx: LoweringContext, op: vector.MultiDimReductionOp\n) -> Sequence[ir.Value]:\n  [in_layout, acc_layout] = inference_utils.in_layouts(op)\n  [out_layout] = inference_utils.out_layouts(op)\n  if out_layout != acc_layout:\n    raise ValueError(\n        f\"Output layout {out_layout} must match the accumulator layout\"\n        f\" {acc_layout}\"\n    )\n\n  if len(op.reduction_dims) != 1:\n    raise NotImplementedError(\"Only 1 reduction dimension is supported.\")\n\n  op_kind = _combining_kind(op.kind)\n  is_signed = _is_reduction_signed(op_kind)\n  src = _fragmented_array_from_ir(op.source, in_layout, is_signed)\n  acc = _fragmented_array_from_ir(op.acc, acc_layout, is_signed)\n\n  if not isinstance(src.layout, fa.TiledLayout):\n    raise NotImplementedError(f\"Unsupported layout: {src.layout}\")\n  reduced_dim = src.layout.tiling.tile_dimension(op.reduction_dims[0])\n  if any(reduced_dim[d] for d in src.layout.partitioned_warp_dims):\n    # cross-warp reductions require scratch space.\n    dtype = op.source.type.element_type\n    allocation_size = ir.IntegerAttr(op.attributes[\"scratch_size\"]).value * 8 // utils.bitwidth(dtype)\n    scratch = _slice_smem(\n        ir.MemRefType.get([allocation_size], dtype, memory_space=utils.smem()),\n        ir.IntegerAttr(op.attributes[\"offset\"]).value,\n        ctx.smem_requested_bytes,\n    )","sourceCodeStart":871,"sourceCodeEnd":907,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/mosaic/gpu/dialect_lowering.py#L871-L907","documentation":"Mosaic only implements multi_dim_reduction over exactly one dimension; reducing 2+ dims at once is unsupported.","triggerScenarios":"vector.multi_dim_reduction with len(op.reduction_dims) != 1.","commonSituations":"Reducing a 3D tile over two axes in one op instead of chaining per-axis reductions.","solutions":["Split into successive single-dimension reductions","Reduce over the tiled dimension last after reshaping to 2D"],"exampleFix":"// before\nr = reduction(v, acc, axes=(1, 2))\n// after\nr = reduction(v, acc, axes=(1,))\nr = reduction(r, acc0, axes=(1,))","handlingStrategy":"validation","validationCode":"assert len(op.reduction_dims) == 1, 'reduce one dim at a time'","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Chain single-axis reductions"],"tags":["jax","mosaic","gpu","reduction"],"backgroundTag":"unsupported-operation","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}