Repository navigation
[CUDA][Performance] Topk is slow #3064
Description
Activity
At sort.cu, we use gpu_sort() which performs a radix sort which could be inefficient(O(n log n)) compared to even our CPU backend where we use std::nth_element which is O(n).
We could use NVIDIA's RAFT library which already has an efficient top-k selection
Reacted by RoyiThe sort in
sort.cuis a merge sort not a radix sort. Also a radix sort is notO(n log n))butO(n)with a constant based on the bit width. For an efficient top-k (esp for half precision) a radix-based selection probably makes sense. Though we will have to play around with that.The sort in
sort.cuis a merge sort not a radix sort. Also a radix sort is notO(n log n))butO(n)with a constant based on the bit width. For an efficient top-k (esp for half precision) a radix-based selection probably makes sense. Though we will have to play around with that.Ah you are right, I was on an earlier commit where we were using CUB's radix sort. The problem then was we were doing full radix sort rather than Radix select. Also if I am not wrong right now we are at
O(n log n))😅 . You are right radix sort isO(w.n).Can I take on this issue?👉🏽👈🏽
Reacted by lin72hFeel free to work on it!
A standalone MLX benchmark for various shape and type combinations:
import time import mlx.core as mx B = [1, 1, 2048, 2048] V = [128_000, 512, 4096, 8192] k = 32 for t in [mx.bfloat16, mx.float32]: for b, v in zip(B, V): q = mx.random.normal(shape=(b, v)).astype(t) def fun(q): for _ in range(50): idx = mx.argpartition(-q, kth=k-1, axis=-1)[:, :k] values = mx.take_along_axis(q, idx, axis=-1) q = mx.put_along_axis(q, idx, values, axis=-1) mx.eval(q) for _ in range(20): fun(q) tic = time.time() for _ in range(20): fun(q) toc = time.time() ms = 1e3 * (toc - tic) print(f"({t}, {b=}, {v=}): {ms=:.3f}")
MLX topk on CUDA is pretty slow in some cases (especially compared to PyTorch).
Here is a benchmark:
On a spark: