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 ctxView on GitHub (pinned to 1e1c6a8fc0)
Solutions
- Unwrap/discard the transform before loading (load from the base memref via unwrap_transformed_memref)
- Use load_tensor/store_tensor ops which do support transformed memrefs
- 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
- Unwrap transforms before scalar loads
- Use load_tensor for transformed buffers
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
- Unsupported transform: {type(transform)}
- Non-indexing transforms on GMEM refs are not implemented.
- Not all transforms could be handled. Remaining transforms: {
- memref.StoreOp does not support transforms: {op}
- Unsupported dtype: {ref.dtype}
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/da55b6bbf01113df.
Report an issue: GitHub.