Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
147 changes: 147 additions & 0 deletions backends/mlx/ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -123,6 +123,7 @@
PartitionNode,
PowerNode,
ProdNode,
PutAlongAxisNode,
RandomBitsNode,
ReciprocalNode,
RemainderNode,
Expand Down Expand Up @@ -2633,6 +2634,152 @@ def _scatter_add_handler(P: MLXProgramBuilder, n: Node) -> Slot:
return out


def _slice_to_shape(
P: MLXProgramBuilder,
x: Slot,
x_shape,
target_shape,
skip_axis: Optional[int],
op_name: str,
) -> Slot:
"""Narrow ``x`` to ``target_shape`` on every axis but ``skip_axis``.

aten.gather and aten.scatter allow the index to be smaller than the
tensor it is applied to (``index.size(d) <= self.size(d)`` for ``d != dim``,
and ``index.size(d) <= src.size(d)`` for every ``d``), and read only the
leading ``index.size(d)`` entries of the larger tensor. MLX broadcasts the
two operands instead, so the larger tensor is sliced down first. A static
target size can always be sliced to; a symbolic one has to be the same
symbol on both sides, since the slice stop is fixed at export time.
"""
for axis, (have, want) in enumerate(zip(x_shape, target_shape)):
if axis == skip_axis:
continue
if isinstance(want, int):
if isinstance(have, int) and have == want:
continue
elif str(have) == str(want):
continue
else:
raise ValueError(
f"{op_name}: symbolic index size must match the input on axis "
f"{axis}, got {have} vs {want}"
)
_, sliced = P.make_tmp_slot()
P.emit(
SliceNode(
x=P.slot_to_tid(x),
out=P.slot_to_tid(sliced),
axis=P.to_int_or_vid(axis),
start=P.to_int_or_vid(0),
stop=P.to_int_or_vid(want),
)
)
x = sliced
return x


@REGISTRY.register(target=[torch.ops.aten.gather.default])
def _gather_handler(P: MLXProgramBuilder, n: Node) -> Slot:
"""Handle aten.gather: out[i][j][k] = x[index[i][j][k]][j][k] for dim=0.

gather(self, dim, index, *, sparse_grad=False) -> Tensor

Maps to mlx::take_along_axis(a, indices, axis). The output takes the shape
of ``index``, so an index that is smaller than ``self`` on the other axes is
handled by narrowing ``self`` to the index shape first.
"""
args = P.args(n)
kwargs = P.kwargs(n)
require_args(args, 3, 3, "aten.gather")
require_kwargs(kwargs, {"sparse_grad"}, "aten.gather")
if kwargs.get("sparse_grad", False):
raise ValueError("aten.gather: sparse_grad=True is not supported")
x, dim, index = args

x_meta = n.args[0].meta.get("val")
index_meta = n.args[2].meta.get("val")
if x_meta is None or index_meta is None:
raise ValueError("aten.gather requires input and index shape metadata")
ndim = len(x_meta.shape)
if ndim == 0 or len(index_meta.shape) != ndim:
raise ValueError(
f"aten.gather: index must have the same non-zero rank as input, got "
f"{len(index_meta.shape)} vs {ndim}"
)
dim = dim % ndim

x = _slice_to_shape(P, x, x_meta.shape, index_meta.shape, dim, "aten.gather")
out = P.make_or_get_slot(n)
P.emit(
TakeAlongAxisNode(
x=P.slot_to_tid(x),
indices=P.slot_to_tid(index),
out=P.slot_to_tid(out),
axis=dim,
)
)
return out


@REGISTRY.register(target=[torch.ops.aten.scatter.src, torch.ops.aten.scatter.value])
def _scatter_handler(P: MLXProgramBuilder, n: Node) -> Slot:
"""Handle aten.scatter: out = self; out[index[i][j]][j] = src[i][j] for dim=0.

scatter.src(self, dim, index, src) -> Tensor
scatter.value(self, dim, index, value) -> Tensor

Maps to mlx::put_along_axis(a, indices, values, axis). ``src`` may be larger
than ``index`` (only its leading ``index.shape`` block is read), so it is
narrowed to the index shape; a scalar ``value`` becomes a 0-D constant that
MLX broadcasts. An index smaller than ``self`` on the non-scatter axes is
handled by the runtime (see exec_put_along_axis). With duplicate indices
along ``dim`` aten leaves the result unspecified; MLX keeps one of the
writes.
"""
args = P.args(n)
require_args(args, 4, 4, "aten.scatter")
require_kwargs(P.kwargs(n), set(), "aten.scatter")
x, dim, index, src = args

x_meta = n.args[0].meta.get("val")
index_meta = n.args[2].meta.get("val")
if x_meta is None or index_meta is None:
raise ValueError("aten.scatter requires input and index shape metadata")
ndim = len(x_meta.shape)
if ndim == 0 or len(index_meta.shape) != ndim:
raise ValueError(
f"aten.scatter: index must have the same non-zero rank as input, got "
f"{len(index_meta.shape)} vs {ndim}"
)
if x_meta.dtype.itemsize == 8:
# mlx ScatterAxis has no GPU kernel for 8-byte element types.
raise ValueError(f"aten.scatter: {x_meta.dtype} input is not supported")
dim = dim % ndim

if isinstance(src, Slot):
src_meta = n.args[3].meta.get("val")
if src_meta is None:
raise ValueError("aten.scatter requires src shape metadata")
src = _slice_to_shape(
P, src, src_meta.shape, index_meta.shape, None, "aten.scatter"
)
else:
src = emit_lifted_constant(P, src, x_meta.dtype)

out = P.make_or_get_slot(n)
P.emit(
PutAlongAxisNode(
x=P.slot_to_tid(x),
indices=P.slot_to_tid(index),
values=P.slot_to_tid(src),
out=P.slot_to_tid(out),
axis=dim,
)
)
return out


@REGISTRY.register(target=[torch.ops.aten.select.int, torch.ops.aten.select_copy.int])
def _select_handler(P: MLXProgramBuilder, n: Node) -> Slot:
"""
Expand Down
34 changes: 34 additions & 0 deletions backends/mlx/runtime/MLXInterpreter.h
Original file line number Diff line number Diff line change
Expand Up @@ -901,6 +901,37 @@ inline void exec_scatter_add(
st.set_tensor(n.out, scatter_add_axis(x, indices, updates, n.axis, s));
}

inline void exec_put_along_axis(
const PutAlongAxisNode& n,
ExecutionState& st,
StreamOrDevice s) {
const auto& x = st.const_tensor_ref(n.x);
const auto& indices = st.const_tensor_ref(n.indices);
const auto& values = st.const_tensor_ref(n.values);
const int rank = static_cast<int>(x.ndim());
int axis = normalize_axis(n.axis, rank, "PutAlongAxis");

// aten.scatter only touches the leading index.shape block of self on the
// non-scatter axes; mlx put_along_axis broadcasts instead, so scatter into
// that block and write it back.
Shape stop = x.shape();
bool narrowed = false;
for (int d = 0; d < rank; ++d) {
if (d != axis && indices.shape(d) != x.shape(d)) {
stop[static_cast<size_t>(d)] = indices.shape(d);
narrowed = true;
}
}
if (!narrowed) {
st.set_tensor(n.out, put_along_axis(x, indices, values, axis, s));
return;
}
Shape start(static_cast<size_t>(rank), 0);
array block =
put_along_axis(slice(x, start, stop, s), indices, values, axis, s);
st.set_tensor(n.out, slice_update(x, block, start, stop, s));
}

inline void
exec_slice(const SliceNode& n, ExecutionState& st, StreamOrDevice s) {
const array& x = st.const_tensor_ref(n.x);
Expand Down Expand Up @@ -2587,6 +2618,9 @@ class Interpreter {
case OpCode::CUMMAX:
ops::exec_cummax(std::get<CummaxNode>(instr.node), st, s);
break;
case OpCode::PUT_ALONG_AXIS:
ops::exec_put_along_axis(std::get<PutAlongAxisNode>(instr.node), st, s);
break;
case OpCode::STACK:
ops::exec_stack(std::get<StackNode>(instr.node), st, s);
break;
Expand Down
13 changes: 12 additions & 1 deletion backends/mlx/serialization/schema.fbs
Original file line number Diff line number Diff line change
Expand Up @@ -472,6 +472,16 @@ table ScatterAddNode {
axis: int32; // Dimension to scatter along
}

// Scatter (overwrite): write values into input at index positions along an axis
// Maps to mlx::put_along_axis(a, indices, values, axis)
table PutAlongAxisNode {
x: Tid (required); // Input tensor to scatter into
indices: Tid (required); // Index tensor (same ndim as x)
values: Tid (required); // Values to write (broadcast to indices)
out: Tid (required);
axis: int32; // Dimension to scatter along
}

table ConcatenateNode {
tensors: [Tid] (required); // List of tensors to concatenate
out: Tid (required);
Expand Down Expand Up @@ -1211,7 +1221,8 @@ union OpNode {
UpdateAndAttendNode,
TruncNode,
FlipNode,
CummaxNode
CummaxNode,
PutAlongAxisNode
// BC: Add new op nodes here (append only)
}

Expand Down
Loading
Loading