{"record":{"id":"4607627b93838d71","repo":"jax-ml/jax","slug":"unhandled-transforms-for-multimem-load-reduce-tr","errorCode":null,"errorMessage":"Unhandled transforms for multimem_load_reduce: {transforms}","messagePattern":"Unhandled transforms for multimem_load_reduce: (.+?)","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/primitives.py","lineNumber":5463,"sourceCode":"    raise RuntimeError(\n        \"Failed to infer the output layout of multimem_load_reduce. Please apply\"\n        \" plgpu.layout_cast to its output right after its creation.\"\n    )\n  if not isinstance(layout, (mgpu.TiledLayout, mgpu.WGStridedFragLayout)):\n    raise ValueError(\n        \"Only tiled and WG strided layouts are supported by\"\n        f\" multimem_load_reduce, but got {layout}\"\n    )\n  dtype = ctx.avals_out[0].dtype\n  transforms = tree.unflatten(transforms_leaves)\n  transform_avals = tree.unflatten(ctx.avals_in[1:])\n  ref_aval = ctx.avals_in[0]\n  assert isinstance(ref_aval, state_types.AbstractRef)\n  ref, _, transforms = lowering._handle_transforms(ctx, ref_aval, ref,\n                                                   transform_avals, transforms,\n                                                   allow_peer_refs=False)\n  if transforms:\n    raise NotImplementedError(\n        f\"Unhandled transforms for multimem_load_reduce: {transforms}\"\n    )\n  multi_ref = ctx.launch_ctx.to_remote_multicast(ref)\n  is_signed = mgpu_utils.is_signed(dtype)\n  arr = mgpu.FragmentedArray.load_reduce_untiled(\n      multi_ref, layout=layout, is_signed=is_signed, reduction=reduction_op\n  )\n  return arr\n\n\n@lowering.register_lowering_rule(multimem_load_reduce_p, mgpu.LoweringSemantics.Warpgroup)\ndef _multimem_load_reduce_lowering_rule_wg(\n    ctx: lowering.LoweringRuleContext, ref, *transforms_leaves, tree, collective_axes, reduction_op,\n):\n  if (mesh_info := ctx.module_ctx.mesh_info) is None:\n    raise ValueError(\n        \"JAX device mesh is required by multimem_load_reduce, but not defined.\"\n    )","sourceCodeStart":5445,"sourceCodeEnd":5481,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/primitives.py#L5445-L5481","documentation":"In the Lane-semantics lowering of multimem_load_reduce, _handle_transforms was invoked with allow_peer_refs=False; if any transforms remain on the ref after handling (e.g. peer-ref transforms), they are unsupported and the lowering aborts.","triggerScenarios":"Passing a ref with residual or peer-related transforms to multimem_load_reduce in a kernel lowered under Lane semantics.","commonSituations":"Cross-shard ref sharing feeding a multimem load; combining multimem ops with experimental ref-transform utilities; regressions after JAX upgrades changing transform semantics.","solutions":["Use a locally-created ref (allocated in this kernel) rather than a peer ref","Strip/avoid extra transforms on the ref before the call","Update JAX to pick up broader transform coverage in _handle_transforms","Fall back to a normal load plus an explicit collective reduce"],"exampleFix":null,"handlingStrategy":"fallback","validationCode":null,"typeGuard":null,"tryCatchPattern":"try:\n    kernel_jit(x)\nexcept NotImplementedError as e:\n    if 'multimem_load_reduce' in str(e):\n        run_load_plus_psum_fallback(x)","preventionTips":["Avoid peer refs and stacked transforms on multimem inputs","Maintain a load+reduce fallback path for transform edge cases"],"tags":["jax","pallas","transforms","multimem","lowering"],"backgroundTag":"unhandled-lowering-transforms","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}