Conversation
aten.gather.default, aten.scatter.src and aten.scatter.value are core ATen, so they reach the partitioner undecomposed and a graph that uses them splits at that node today. gather lowers onto the existing TakeAlongAxisNode; the input is first narrowed to the index shape on the non-gather axes, because aten reads only that block while mlx broadcasts. scatter gets a PutAlongAxisNode, appended at the end of the OpNode union, whose interpreter kernel calls mlx put_along_axis; when the index is smaller than self on the other axes it scatters into the leading index-shaped block of self and writes it back with slice_update. A src larger than the index is narrowed the same way and a scalar value becomes a 0-D constant. 8-byte element types are declined since mlx ScatterAxis has no GPU kernel for them. 32 configurations in test_ops.py (1D-4D, every axis, negative dim, index smaller than input and longer along dim, src larger than index, scalar value, fp16/bf16/int inputs, gather from a transposed view, dynamic batch run at an unseen size) match eager exactly through the C++ op_test_runner; the full MLX op suite passes 982/982.
🔗 Helpful Links🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/executorch/23398
Note: Links to docs will display an error until the docs builds have been completed. ❌ 16 Awaiting Approval, 1 New FailureAs of commit e6b6332 with merge base 53c79ed ( AWAITING APPROVAL - The following workflows need approval before CI can run:
NEW FAILURE - The following job has failed:
This comment was automatically generated by Dr. CI and updates every 15 minutes. |
|
Hi @zsun6! Thank you for your pull request and welcome to our community. Action RequiredIn order to merge any pull request (code, docs, etc.), we require contributors to sign our Contributor License Agreement, and we don't seem to have one on file for you. ProcessIn order for us to review and merge your suggested changes, please sign at https://code.facebook.com/cla. If you are contributing on behalf of someone else (eg your employer), the individual CLA may not be sufficient and your employer may need to sign the corporate CLA. Once the CLA is signed, our tooling will perform checks and validations. Afterwards, the pull request will be tagged with If you have received this in error or have any questions, please contact us at cla@meta.com. Thanks! |
|
|
This PR needs a
|
aten.gather.default,aten.scatter.srcandaten.scatter.valuehad no MLX handler. They're core ATen soto_edgedoesn't decompose them, and any graph with one of them splits at that node: the op runs on the portable CPU kernel with a delegate handoff on each side.Subgraphs from this repo, before/after (
to_edge_transform_and_lowerwithMLXPartitioner):rel_shift(examples/models/sortformer/export_sortformer.py:164, one per attention layer), x (1,8,64,127)gatheron CPUsampling_probabilities(examples/models/muse-glimmer/model/dflash_token_sampler.py:63), logits (4,512)scatter.src,sort.stableon CPUsort.stableon CPUverify_speculative(same file, R=3, V=512)gather,scatter.src,scatter.value,sort.stable, 2xsumon CPUsort.stable, 2xsumon CPUexamples/models/llama/llama_transformer.py:156(scores.gather(1, idx)) and HF gpt_oss / qwen3_vl_moe (zeros_like(logits).scatter_(1, idx, w)), (16,256) -> 8 expertsgather,scatter.srcon CPU(
sort.stableisn't core ATen and needs_core_aten_ops_exception_list;sumover bool is a separate gap. Both out of scope here.)How
gatherlowers onto the existingTakeAlongAxisNode. aten allowsindex.size(d) <= self.size(d)on the non-gather axes and only reads that leading block, whilemlx::take_along_axisbroadcasts, so the handler first narrowsselfto the index shape on those axes withSliceNodes (nothing is emitted when the sizes already agree, including a shared symbolic dim; a static index size is sliced to even if the input dim is symbolic). The output has the index shape, so the index can be longer thanselfalongdimfor free.scatteradds aPutAlongAxisNode(appended at the end of theOpNodeunion, so existing.ptefiles still load) and anexec_put_along_axisthat callsmlx::put_along_axis. Two aten freedoms need handling:srcmay be larger thanindexon any axis (only its leadingindex.shapeblock is read), which the handler narrows like gather; andindexmay be smaller thanselfon the non-scatter axes, which the kernel handles by scattering into the leading index-shaped block ofselfand writing it back withslice_update— broadcasting there would write every row instead of the firstindex.size(d). A scalarvaluebecomes a 0-D constant that MLX broadcasts, soscatter.valueshares the node.int64/float64inputs are declined at partition time becauseScatterAxis::eval_gpuhas no kernel for 8-byte element types. Duplicate indices alongdimare unspecified in aten; MLX keeps one of the writes. Indices aren't bounds-checked in the delegate, same as the existingtake_along_axis/scatter_addlowerings.Tests
Adds
gatherandscattertobackends/mlx/test/test_ops.py, 32 configs: 1D-4D, every axis and negativedim, index smaller than the input on the other axes and longer alongdim,srclarger thanindex, scalarvalue(including a float written into an int32self), fp16/bf16/int64 inputs with int64 indices, a gather from a transposed view, and a dynamic batch dim (shared by the operands, and on the input alone against a static index) exported at one size and run at another. Both ops are pure data movement, so the comparison is exact (rtol = atol = 0).M3 Pro, macOS 14.1, Xcode 15.3, torch 2.14.0, MLX 1f8e74e3. Only
schema.fbsis committed; the generated files come from the build, as in #23051.Written with AI assistance (Claude Code). I reviewed the design and the diff, ran the full MLX op suite locally and take responsibility for the change.
cc @nil-is-all @metascroy