jax-ml/jax · error · NotImplementedError

memref.LoadOp does not support transforms: {op}

Error message

memref.LoadOp does not support transforms: {op}

What it means

memref.load is a scalar access and Mosaic's lowering is a pass-through: it explicitly rejects any non-empty transform annotation on the loaded memref, since transformed layouts can't be indexed directly.

Source

Thrown at jax/experimental/mosaic/gpu/dialect_lowering.py:2363

      unwrap_transformed_memref(op.src, in_transforms_attr),
      new_reassociation,
  )
  return [wrap_transformed_memref(result, op.result.type, out_transforms_attr)]


@_register_lowering(memref.LoadOp)
def _memref_load_op_lowering_rule(
    ctx: LoweringContext, op: memref.LoadOp
) -> Sequence[ir.Value]:
  """Lowering rule for memref.LoadOp.

  Loads are never transformed so this rule is mostly just a pass-through.
  """
  del ctx

  in_transforms = inference_utils.in_transforms(op)[0]
  if in_transforms:
    raise NotImplementedError(f"memref.LoadOp does not support transforms: {op}")

  new_load_op = memref.LoadOp(
      memref=unwrap_transformed_memref(op.memref, in_transforms),
      indices=op.indices,
      nontemporal=op.nontemporal,
  )
  return [new_load_op.result]


@_register_lowering(memref.StoreOp, support_warp_semantics=True)
def _memref_store_op_lowering_rule(
    ctx: LoweringContext, op: memref.StoreOp
) -> Sequence[ir.Value]:
  """Lowering rule for memref.StoreOp.

  Stores are never transformed so this rule is mostly just a pass-through.
  """
  del ctx

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Unwrap/discard the transform before loading (load from the base memref via unwrap_transformed_memref)
  2. Use load_tensor/store_tensor ops which do support transformed memrefs
  3. Avoid tiling memrefs you intend to access element-wise

Example fix

// before
v = t.memref.load(tiled_ref, idx)
// after
base = unwrap_transformed_memref(tiled_ref, transforms)
v = t.memref.load(base, idx)
Defensive patterns

Strategy: validation

Validate before calling

assert not inference_utils.in_transforms(op)[0], 'load does not support transforms'

Type guard

def loadable_untransformed(op) -> bool:
    return not inference_utils.in_transforms(op)[0]

Prevention

When it happens

Trigger: Emitting memref.load where inference_utils.in_transforms(op)[0] is non-empty, i.e. loading from a memref that still carries tiling/swizzle transforms.

Common situations: Loading individual elements from a tiled/swizzled smem tensor without unwrapping the transform first; using low-level load instead of Mosaic's tensor loads on transformed buffers.

Related errors


AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27). Data as JSON: /api/errors/da55b6bbf01113df. Report an issue: GitHub.