Skip to content

[Metal][Performance] Add radix select kernel for partition and argpartition - #4559

Closed
mateuuszzzzz wants to merge 4 commits into
ml-explore:mainfrom
mateuuszzzzz:add-metal-radix-partition
Closed

mateuuszzzzz wants to merge 4 commits into
ml-explore:mainfrom
mateuuszzzzz:add-metal-radix-partition

Conversation

@mateuuszzzzz

@mateuuszzzzz mateuuszzzzz commented Sep 24, 2026 •

Copy link
Copy Markdown
Contributor

Details

Currently mx.partition and mx.argpartition on Metal are routed straight to
the merge sort kernel (gpu_merge_sort), so a partition costs as much as a full
sort of the row. For most shapes this is inefficient: a radix select only needs
to find the k-th key and split the row around it, which speeds these operations
up by as much as 17x for certain shapes. This PR proposes a dedicated radix select kernel and
dispatches long rows to it, keeping the merge sort for everything else. Addressed #3064 (previous #3069 was closed)

A few measured choices in the kernel

  • Digit width follows the threadgroup size: 11 bits (2048 bins) at 256 threads, 1.06 to 1.34x over 8 bits on f32 from fewer passes over the row. 64 and 128 threads use 8 bits.
  • Threadgroup size 64, 128 or 256 by row length: below 1024 columns, below 4096 columns, the rest.
  • Candidates cache: once the keys matching the prefix fit in threadgroup memory (8 per thread), later passes read the cache instead of the row, 1.11 to 1.66x over no cache.
  • Counting: 16 loads in flight per thread, 1.5 to 1.74x over one load at a time
  • Write: 8 consecutive elements per thread and one threadgroup scan per chunk, which positions below, equal and above at once and keeps input order, so ties come out like the sort. 8 per thread is 6 to 10% faster than 4, 16 brings nothing more. Full chunks skip bounds checks, +5 to 14% at low row counts.
  • Histogram in threadgroup memory on relaxed atomics with barriers between phases, so the whole partition is a single dispatch with no temporary buffers.

Benchmarks

benchmarks/python/partition_bench.py on an M3 Pro 18 GB, main 59d600b5e vs this PR; partition with argpartition in brackets

shape k float32 main [ms] float32 PR [ms] speedup bfloat16 main [ms] bfloat16 PR [ms] speedup
4096×8192 2048 22.64 (22.65) 2.43 (2.35) 9.3x (9.6x) 19.63 (20.67) 1.59 (1.87) 12.3x (11.0x)
2048×32768 2048 63.18 (65.21) 6.66 (6.77) 9.5x (9.6x) 53.38 (57.72) 3.20 (4.18) 16.7x (13.8x)
2048×8192 32 11.43 (11.52) 1.31 (1.30) 8.7x (8.9x) 9.98 (10.45) 0.834 (1.26) 12.0x (8.3x)
64×128000 40 10.17 (10.24) 1.31 (1.29) 7.8x (7.9x) 8.18 (8.39) 0.726 (1.06) 11.3x (8.0x)
8×131072 2048 1.45 (1.37) 0.337 (0.338) 4.3x (4.0x) 1.17 (1.20) 0.275 (0.306) 4.3x (3.9x)
2048×1024 32 0.744 (0.813) 0.310 (0.315) 2.4x (2.6x) 0.741 (0.826) 0.265 (0.275) 2.8x (3.0x)
4096×256 8 0.340 (0.372) 0.282 (0.276) 1.2x (1.3x) 0.341 (0.389) 0.251 (0.368) 1.4x (1.1x)
32×4096 32 0.217 (0.217) 0.132 (0.134) 1.6x (1.6x) 0.218 (0.208) 0.107 (0.121) 2.0x (1.7x)
Rows on both sides of the dispatch thresholds (129 columns, 256K elements, 1024 columns, 4096 columns)

Note that the 0.9x results correspond to cases where the existing sort kernel is selected. This difference is within measurement noise, so these cases should be considered performance-neutral rather than regressions. The meaningful gains are on shapes where the new kernel is selected.

shape k float32 main [ms] float32 PR [ms] speedup bfloat16 main [ms] bfloat16 PR [ms] speedup
8192×128 8 0.312 (0.335) 0.311 (0.335) 1.0x (1.0x) 0.305 (0.355) 0.341 (0.393) 0.9x (0.9x)
8192×129 8 0.535 (0.614) 0.360 (0.364) 1.5x (1.7x) 0.546 (0.642) 0.343 (0.352) 1.6x (1.8x)
511×512 32 0.186 (0.195) 0.186 (0.197) 1.0x (1.0x) 0.192 (0.200) 0.189 (0.188) 1.0x (1.1x)
512×512 32 0.169 (0.163) 0.139 (0.139) 1.2x (1.2x) 0.184 (0.187) 0.117 (0.139) 1.6x (1.3x)
2048×1023 32 0.677 (0.756) 0.304 (0.306) 2.2x (2.5x) 0.729 (0.823) 0.283 (0.274) 2.6x (3.0x)
32×4095 32 0.218 (0.218) 0.215 (0.215) 1.0x (1.0x) 0.215 (0.218) 0.188 (0.204) 1.1x (1.1x)

Limitations of current PR

  • There are two cases where this kernel does not win:

    • a few very long rows (1 to 8 × 32K to 128K), where a single threadgroup leaves most of the GPU idle,
    • short rows (up to 512 columns), where zeroing and scanning the histogram costs about as much as counting the row

    If this PR is accepted, I would like to address both with dedicated follow-up kernels rather than here, since this PR is already large. The current dispatch rule only routes shapes that benefit from the radix kernel to it and leaves the rest on the merge sort path, so there is no performance regression in the meantime.

  • Non-contiguous inputs. The radix path now supports contiguous inputs with the partition axis at stride 1, while other layouts stay on the merge sort. A copy to a contiguous buffer before the kernel was measured too: it gives a median of 2.2x and up to 8.7x over the sort on views, with the only flat area being single rows. It is left out of this PR because on a single row it buys nothing until there is a kernel for few-row shapes. If there is interest in the follow-up kernels for few-row and short-row shapes, enabling the copy then would make the radix path a win for every shape.


  • ☑️ I understand it is strictly prohibited to use AI to write PR description
  • AI usage disclosure: I used AI as coding assistant, for alignment with codebase conventions, for generating test cases and finding various edge cases

@mateuuszzzzz

mateuuszzzzz commented Sep 24, 2026 •

Copy link
Copy Markdown
Contributor Author

Also, I ran into a couple of difficulties during development that may be standalone issues:

complex64_t ordering with NaN: The radix key reproduces the ordering of the Metal sort (real part first, then imaginary, with NaN in either part last). The CPU sort currently orders these cases differently. The complex64_t test always compares against the sort on the same device, but the NaN cases are therefore limited to GPU. #4519 aligns the CPU ordering with Metal. Once it is merged RadixKeyTraits<complex64_t> needs to be aligned.

bool: There is no bool instantiation because sort.metal has no bool kernel either. Sort, partition, and topk on bool already fail on main with a missing kernel error on Metal, while CPU and CUDA work. If someone ever adds bool support for sort, adding a bool instantiation here is a one-line change, the generic trait already handles bool as uint8 for partition. I'm not sure whether the lack of bool support in sort is intentional.

@mateuuszzzzz mateuuszzzzz changed the title Add a Metal radix select kernel for partition and argpartition [Metal][Performance] Add radix select kernel for partition and argpartition Sep 24, 2026
@zcbenz

zcbenz commented Sep 25, 2026

Copy link
Copy Markdown
Member

There is a limitation on contributor PRs because we are unable to spend too much time reviewing contributor PRs at the moment, please do not open new PRs before we review your current open ones.

Regarding to this PR, it is not hard to write a radix sort manually or ask AI to do that, but you would need to know a lot more to optimize for the cases we care.

@zcbenz zcbenz closed this Sep 25, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants