Skip to content

[CUDA][Performance] Topk is slow #3064

Description

@awni

MLX topk on CUDA is pretty slow in some cases (especially compared to PyTorch).

Here is a benchmark:

import time
import mlx.core as mx

b = 2048
v = 8192
k = 32

q = mx.random.normal(shape=(b, v)).astype(mx.bfloat16)

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"MLX {ms=:.3f}")

import torch


q = torch.randn(size=(b, v)).to("cuda").to(torch.bfloat16)
def topk_old(q):
    return idx, values


def fun(q):
    for _ in range(50):
        values, idx = torch.topk(q, k=k, axis=-1)
        q = torch.scatter(q, -1, idx, values)
    torch.cuda.synchronize()

for _ in range(20):
    fun(q)

tic = time.time()
for _ in range(20):
    fun(q)
toc = time.time()
ms = 1e3 * (toc - tic)
print(f"PyTorch {ms=:.3f}")

On a spark:

MLX ms=2975.014
PyTorch ms=919.764

Activity

  1. NripeshN commented on Jan 26, 2026

    @NripeshN
    Contributor

    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

  2. awni commented on Jan 26, 2026

    @awni
    MemberAuthor

    The sort in sort.cu is a merge sort not a radix sort. Also a radix sort is not O(n log n)) but O(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.

  3. NripeshN commented on Jan 27, 2026

    @NripeshN
    Contributor

    The sort in sort.cu is a merge sort not a radix sort. Also a radix sort is not O(n log n)) but O(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 is O(w.n).

    Can I take on this issue?👉🏽👈🏽

  4. awni commented on Jan 27, 2026

    @awni
    MemberAuthor

    Feel free to work on it!

  5. awni commented on Jan 27, 2026

    @awni
    MemberAuthor

    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}")
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions