[Metal][Performance] Add radix select kernel for partition and argpartition - #4559
mateuuszzzzz wants to merge 4 commits into
Conversation
|
Also, I ran into a couple of difficulties during development that may be standalone issues:
|
partition and argpartitionpartition and argpartition
|
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. |
Details
Currently
mx.partitionandmx.argpartitionon Metal are routed straight tothe merge sort kernel (
gpu_merge_sort), so a partition costs as much as a fullsort 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
Benchmarks
benchmarks/python/partition_bench.pyon an M3 Pro 18 GB, main59d600b5evs this PR;partitionwithargpartitionin bracketsRows 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.
Limitations of current PR
There are two cases where this kernel does not win:
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.