diff --git a/backends/mlx/ops.py b/backends/mlx/ops.py index 3973f1c80d7..9885d5c59c9 100644 --- a/backends/mlx/ops.py +++ b/backends/mlx/ops.py @@ -123,6 +123,7 @@ PartitionNode, PowerNode, ProdNode, + PutAlongAxisNode, RandomBitsNode, ReciprocalNode, RemainderNode, @@ -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: """ diff --git a/backends/mlx/runtime/MLXInterpreter.h b/backends/mlx/runtime/MLXInterpreter.h index 4c68bd40022..99dc2ed587d 100644 --- a/backends/mlx/runtime/MLXInterpreter.h +++ b/backends/mlx/runtime/MLXInterpreter.h @@ -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(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(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(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); @@ -2587,6 +2618,9 @@ class Interpreter { case OpCode::CUMMAX: ops::exec_cummax(std::get(instr.node), st, s); break; + case OpCode::PUT_ALONG_AXIS: + ops::exec_put_along_axis(std::get(instr.node), st, s); + break; case OpCode::STACK: ops::exec_stack(std::get(instr.node), st, s); break; diff --git a/backends/mlx/serialization/schema.fbs b/backends/mlx/serialization/schema.fbs index 2743c093a17..95354e9c045 100644 --- a/backends/mlx/serialization/schema.fbs +++ b/backends/mlx/serialization/schema.fbs @@ -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); @@ -1211,7 +1221,8 @@ union OpNode { UpdateAndAttendNode, TruncNode, FlipNode, - CummaxNode + CummaxNode, + PutAlongAxisNode // BC: Add new op nodes here (append only) } diff --git a/backends/mlx/test/test_ops.py b/backends/mlx/test/test_ops.py index 6942cd7ae1b..310520430cd 100644 --- a/backends/mlx/test/test_ops.py +++ b/backends/mlx/test/test_ops.py @@ -7403,6 +7403,256 @@ def create_inputs(self) -> Tuple[torch.Tensor, ...]: return (x, index, src) +class GatherModel(nn.Module): + """Model that gathers along a dimension, optionally from a transposed view.""" + + def __init__(self, dim: int = 0, transpose: bool = False): + super().__init__() + self.dim = dim + self.transpose = transpose + + def forward(self, x: torch.Tensor, index: torch.Tensor) -> torch.Tensor: + if self.transpose: + x = x.transpose(0, 1) + return torch.gather(x, self.dim, index) + + +@register_test +class GatherTest(OpTestCase): + """Test case for aten.gather. + + gather(self, dim, index) reads self at index positions along dim. The + index may be longer than self along dim and smaller than self on the + other axes; the output takes the index shape. Pure data movement, so the + comparison is exact. + """ + + name = "gather" + rtol = 0 + atol = 0 + + def __init__( + self, + shape: Tuple[int, ...] = (4, 8), + dim: int = 1, + index_shape: Optional[Tuple[int, ...]] = None, + dtype: torch.dtype = torch.float32, + transpose: bool = False, + dynamic_batch: bool = False, + ): + self.shape = shape + self.dim = dim + self.index_shape = index_shape + self.dtype = dtype + self.transpose = transpose + self.dynamic_batch = dynamic_batch + parts = ["gather", "x".join(str(s) for s in shape), f"dim{dim}"] + if index_shape is not None: + parts.append("idx" + "x".join(str(s) for s in index_shape)) + if dtype != torch.float32: + parts.append(str(dtype).replace("torch.", "")) + if transpose: + parts.append("t") + if dynamic_batch: + parts.append("dyn") + self.name = "_".join(parts) + + @classmethod + def get_test_configs(cls) -> List["GatherTest"]: + return [ + # 1D, index longer than input along dim + cls(shape=(8,), dim=0, index_shape=(12,)), + # 2D, each axis, negative dim + cls(shape=(4, 8), dim=0), + cls(shape=(4, 8), dim=1), + cls(shape=(4, 8), dim=-1), + # index smaller than input on the non-gather axis + cls(shape=(4, 8), dim=1, index_shape=(2, 5)), + # 3D: longer along dim, smaller on the other axes + cls(shape=(2, 4, 8), dim=1, index_shape=(1, 6, 3)), + cls(shape=(2, 4, 8), dim=0), + # 4D + cls(shape=(2, 3, 4, 8), dim=-1), + cls(shape=(2, 3, 4, 8), dim=2, index_shape=(2, 3, 2, 8)), + # half precision and integer inputs + cls(shape=(4, 8), dim=1, dtype=torch.float16), + cls(shape=(4, 8), dim=0, dtype=torch.bfloat16), + cls(shape=(4, 8), dim=1, dtype=torch.int64), + # gather from a transposed (non-contiguous) view + cls(shape=(4, 8), dim=1, transpose=True), + # dynamic batch shared by input and index, and on the input alone + # with a static index narrowing it; both run at an unseen batch + cls(shape=(4, 8), dim=1, dynamic_batch=True), + cls(shape=(4, 8), dim=1, index_shape=(2, 5), dynamic_batch=True), + ] + + def create_model(self) -> nn.Module: + return GatherModel(dim=self.dim, transpose=self.transpose) + + def get_dynamic_shapes(self) -> Optional[Dict[str, any]]: + if not self.dynamic_batch: + return None + batch = Dim("batch", min=2, max=16) + return {"x": {0: batch}, "index": None if self.index_shape else {0: batch}} + + def create_inputs(self) -> Tuple[torch.Tensor, ...]: + return self._make_inputs(self.shape) + + def create_test_inputs(self) -> Tuple[torch.Tensor, ...]: + if not self.dynamic_batch: + return self.create_inputs() + return self._make_inputs((self.shape[0] + 3,) + self.shape[1:]) + + def _make_inputs(self, shape: Tuple[int, ...]) -> Tuple[torch.Tensor, ...]: + if self.dtype.is_floating_point: + x = torch.randn(shape, dtype=self.dtype) + else: + x = torch.randint(-8, 8, shape, dtype=self.dtype) + view_shape = list(shape) + if self.transpose: + view_shape[0], view_shape[1] = view_shape[1], view_shape[0] + index_shape = self.index_shape or tuple(view_shape) + index = torch.randint(0, view_shape[self.dim], index_shape, dtype=torch.long) + return (x, index) + + +class ScatterModel(nn.Module): + """Model that scatters src (or a scalar value) into x along a dimension.""" + + def __init__(self, dim: int = 0, value: Optional[float] = None): + super().__init__() + self.dim = dim + self.value = value + + def forward(self, x: torch.Tensor, index: torch.Tensor, src: torch.Tensor): + if self.value is not None: + return x.scatter(self.dim, index, self.value) + return x.scatter(self.dim, index, src) + + +@register_test +class ScatterTest(OpTestCase): + """Test case for aten.scatter.src and aten.scatter.value. + + scatter(self, dim, index, src) writes src at index positions along dim. + The index may be smaller than self on the other axes and smaller than + src on every axis. Indices are unique along dim so the result is + well-defined, and the comparison is exact. + """ + + name = "scatter" + rtol = 0 + atol = 0 + + def __init__( + self, + shape: Tuple[int, ...] = (4, 8), + dim: int = 1, + num_indices: int = 3, + index_shape: Optional[Tuple[int, ...]] = None, + src_shape: Optional[Tuple[int, ...]] = None, + value: Optional[float] = None, + dtype: torch.dtype = torch.float32, + dynamic_batch: bool = False, + ): + self.shape = shape + self.dim = dim + self.num_indices = num_indices + self.index_shape = index_shape + self.src_shape = src_shape + self.value = value + self.dtype = dtype + self.dynamic_batch = dynamic_batch + parts = ["scatter", "x".join(str(s) for s in shape), f"dim{dim}"] + parts.append(f"idx{num_indices}") + if index_shape is not None: + parts.append("ishape" + "x".join(str(s) for s in index_shape)) + if src_shape is not None: + parts.append("src" + "x".join(str(s) for s in src_shape)) + if value is not None: + parts.append("value") + if dtype != torch.float32: + parts.append(str(dtype).replace("torch.", "")) + if dynamic_batch: + parts.append("dyn") + self.name = "_".join(parts) + + @classmethod + def get_test_configs(cls) -> List["ScatterTest"]: + return [ + # 1D + cls(shape=(8,), dim=0, num_indices=5), + # 2D, each axis, negative dim, full permutation along dim + cls(shape=(4, 8), dim=1, num_indices=3), + cls(shape=(4, 8), dim=0, num_indices=2), + cls(shape=(4, 8), dim=-1, num_indices=8), + # 3D and 4D + cls(shape=(2, 4, 8), dim=1, num_indices=2), + cls(shape=(2, 3, 4, 8), dim=-1, num_indices=4), + # index smaller than self on the non-scatter axis + cls(shape=(4, 8), dim=1, num_indices=3, index_shape=(2, 3)), + cls(shape=(2, 4, 8), dim=0, num_indices=1, index_shape=(1, 3, 5)), + # src larger than index + cls(shape=(4, 8), dim=1, num_indices=3, src_shape=(4, 6)), + cls(shape=(4, 8), dim=0, num_indices=2, src_shape=(3, 8)), + # scalar value (aten.scatter.value) + cls(shape=(4, 8), dim=1, num_indices=3, value=2.5), + cls(shape=(2, 4, 8), dim=0, num_indices=1, value=-1.0), + # half precision, and a float value cast into an int32 self + cls(shape=(4, 8), dim=1, num_indices=3, dtype=torch.float16), + cls(shape=(4, 8), dim=0, num_indices=2, dtype=torch.bfloat16), + cls(shape=(4, 8), dim=1, num_indices=3, value=2.5, dtype=torch.int32), + # dynamic batch shared by all three inputs, and on self alone with + # a static index narrowing it at runtime; both run at an unseen batch + cls(shape=(4, 8), dim=1, num_indices=3, dynamic_batch=True), + cls( + shape=(4, 8), + dim=1, + num_indices=3, + index_shape=(2, 3), + dynamic_batch=True, + ), + ] + + def create_model(self) -> nn.Module: + return ScatterModel(dim=self.dim, value=self.value) + + def get_dynamic_shapes(self) -> Optional[Dict[str, any]]: + if not self.dynamic_batch: + return None + batch = Dim("batch", min=2, max=16) + rest = None if self.index_shape else {0: batch} + return {"x": {0: batch}, "index": rest, "src": rest} + + def create_inputs(self) -> Tuple[torch.Tensor, ...]: + return self._make_inputs(self.shape) + + def create_test_inputs(self) -> Tuple[torch.Tensor, ...]: + if not self.dynamic_batch: + return self.create_inputs() + return self._make_inputs((self.shape[0] + 3,) + self.shape[1:]) + + def _make_inputs(self, shape: Tuple[int, ...]) -> Tuple[torch.Tensor, ...]: + def rand(s): + if self.dtype.is_floating_point: + return torch.randn(s, dtype=self.dtype) + return torch.randint(-8, 8, s, dtype=self.dtype) + + x = rand(shape) + dim = self.dim % len(shape) + index_shape = list(self.index_shape or shape) + index_shape[dim] = self.num_indices + # A random permutation along dim, cut to num_indices, gives unique + # targets per output position. + perm_shape = list(index_shape) + perm_shape[dim] = shape[dim] + index = ( + torch.rand(perm_shape).argsort(dim).narrow(dim, 0, self.num_indices) + ).contiguous() + src = rand(self.src_shape or index_shape) + return (x, index, src) + + @register_test class QuantizedEmbeddingTest(OpTestCase): """Test case for TorchAO int4 quantized nn.Embedding."""