Skip to content

MLX: add gather and scatter handlers - #23398

Open
zsun6 wants to merge 1 commit into
pytorch:mainfrom
zsun6:feat/mlx-gather-scatter
Open

zsun6 wants to merge 1 commit into
pytorch:mainfrom
zsun6:feat/mlx-gather-scatter

Conversation

@zsun6

@zsun6 zsun6 commented Oct 3, 2026 •

Copy link
Copy Markdown

aten.gather.default, aten.scatter.src and aten.scatter.value had no MLX handler. They're core ATen so to_edge doesn'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_lower with MLXPartitioner):

subgraph before after
Sortformer export-safe rel_shift (examples/models/sortformer/export_sortformer.py:164, one per attention layer), x (1,8,64,127) 1 segment, gather on CPU 1 segment, nothing on CPU
Muse Glimmer sampling_probabilities (examples/models/muse-glimmer/model/dflash_token_sampler.py:63), logits (4,512) 3 segments; scatter.src, sort.stable on CPU 2 segments; sort.stable on CPU
Muse Glimmer verify_speculative (same file, R=3, V=512) 5 segments; 2x gather, scatter.src, scatter.value, sort.stable, 2x sum on CPU 3 segments; sort.stable, 2x sum on CPU
top-2 MoE router as in examples/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 experts 2 segments; gather, scatter.src on CPU 1 segment, nothing on CPU

(sort.stable isn't core ATen and needs _core_aten_ops_exception_list; sum over bool is a separate gap. Both out of scope here.)

How

gather lowers onto the existing TakeAlongAxisNode. aten allows index.size(d) <= self.size(d) on the non-gather axes and only reads that leading block, while mlx::take_along_axis broadcasts, so the handler first narrows self to the index shape on those axes with SliceNodes (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 than self along dim for free.

scatter adds a PutAlongAxisNode (appended at the end of the OpNode union, so existing .pte files still load) and an exec_put_along_axis that calls mlx::put_along_axis. Two aten freedoms need handling: src may be larger than index on any axis (only its leading index.shape block is read), which the handler narrows like gather; and index may be smaller than self on the non-scatter axes, which the kernel handles by scattering into the leading index-shaped block of self and writing it back with slice_update — broadcasting there would write every row instead of the first index.size(d). A scalar value becomes a 0-D constant that MLX broadcasts, so scatter.value shares the node. int64/float64 inputs are declined at partition time because ScatterAxis::eval_gpu has no kernel for 8-byte element types. Duplicate indices along dim are unspecified in aten; MLX keeps one of the writes. Indices aren't bounds-checked in the delegate, same as the existing take_along_axis/scatter_add lowerings.

Tests

Adds gather and scatter to backends/mlx/test/test_ops.py, 32 configs: 1D-4D, every axis and negative dim, index smaller than the input on the other axes and longer along dim, src larger than index, scalar value (including a float written into an int32 self), 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).

python -m executorch.backends.mlx.test.run_all_tests --rebuild gather scatter   # 32 passed, 0 failed
python -m executorch.backends.mlx.test.run_all_tests -j4 --clean-after           # 982 passed, 0 failed

M3 Pro, macOS 14.1, Xcode 15.3, torch 2.14.0, MLX 1f8e74e3. Only schema.fbs is 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

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.
@pytorch-bot

pytorch-bot Bot commented Oct 3, 2026 •

Copy link
Copy Markdown

🔗 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 Failure

As of commit e6b6332 with merge base 53c79ed (image):

AWAITING APPROVAL - The following workflows need approval before CI can run:

NEW FAILURE - The following job has failed:

  • Cadence Build & Test / Resolve CI docker image / resolve (gh)
    ##[error]Refusing to check out fork pull request code from a 'pull_request_target' workflow. This workflow runs with the base repository's GITHUB_TOKEN, secrets, default-branch cache scope, and runner access. Fetching and executing a fork's code in that trusted context commonly leads to "pwn request" vulnerabilities. To opt in, review the risks at https://gh.io/securely-using-pull_request_target and set 'allow-unsafe-pr-checkout: true' on the actions/checkout step.

This comment was automatically generated by Dr. CI and updates every 15 minutes.

@meta-cla

meta-cla Bot commented Oct 3, 2026

Copy link
Copy Markdown

Hi @zsun6!

Thank you for your pull request and welcome to our community.

Action Required

In 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.

Process

In 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 CLA signed. The tagging process may take up to 1 hour after signing. Please give it that time before contacting us about it.

If you have received this in error or have any questions, please contact us at cla@meta.com. Thanks!

@linux-foundation-easycla

linux-foundation-easycla Bot commented Oct 3, 2026 •

Copy link
Copy Markdown

CLA Signed
The committers listed above are authorized under a signed CLA.

  • ✅ login: zsun6 / name: Zhongrui Sun (e6b6332)

@github-actions

github-actions Bot commented Oct 3, 2026

Copy link
Copy Markdown

This PR needs a release notes: label

If your change should be included in the release notes (i.e. would users of this library care about this change?), please use a label starting with release notes:. This helps us keep track and include your important work in the next release notes.

To add a label, you can comment to pytorchbot, for example
@pytorchbot label "release notes: none"

For more information, see
https://github.com/pytorch/pytorch/wiki/PyTorch-AutoLabel-Bot#why-categorize-for-release-notes-and-how-does-it-work.

@executorch-triage executorch-triage Bot added the community: contribution PRs coming from community (excluding hardware partners) label Oct 3, 2026

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

community: contribution PRs coming from community (excluding hardware partners)

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant