Skip to content

ROCm: resolve TOP_K kernels - #28313

Open
pwilkin wants to merge 3 commits into
ggml-org:masterfrom
pwilkin:topk-rocm-fix
Open

pwilkin wants to merge 3 commits into
ggml-org:masterfrom
pwilkin:topk-rocm-fix

Conversation

@pwilkin

@pwilkin pwilkin commented Sep 3, 2026 •

Copy link
Copy Markdown
Member

Overview

The current TOP_K kernels are still far from optimal, @Geramy 's native library fix resolves some of the cases, but not all of them.

Additional information

Added a small-case kernel (ported @0cc4m 's Vulkan kernel basically) and sorted out the optimal cases. Comparison table (tests disabled HIP graphs because the graph update algorithm is currently broken in upstream code pending ROCm/rocm-systems#11069):

# TOP_K parameters Master us Fix us Speedup Time reduction
1 type=f32,ne=[2,1,1,1],k=1,ties=0 4.11 2.36 1.742x 42.58%
2 type=f32,ne=[4096,1,1,1],k=16,ties=0 43.56 8.39 5.192x 80.74%
3 type=f32,ne=[4096,16,1,1],k=16,ties=0 45.30 12.16 3.725x 73.16%
4 type=f32,ne=[8192,1,1,1],k=16,ties=0 38.58 9.30 4.148x 75.89%
5 type=f32,ne=[8192,16,1,1],k=16,ties=0 42.35 19.41 2.182x 54.17%
6 type=f32,ne=[12288,1,1,1],k=16,ties=0 43.19 10.68 4.044x 75.27%
7 type=f32,ne=[12288,16,1,1],k=16,ties=0 51.64 25.92 1.992x 49.81%
8 type=f32,ne=[16384,1,1,1],k=16,ties=0 42.90 11.37 3.773x 73.50%
9 type=f32,ne=[16384,16,1,1],k=16,ties=0 52.86 31.07 1.701x 41.22%
10 type=f32,ne=[24576,1,1,1],k=16,ties=0 43.76 14.56 3.005x 66.73%
11 type=f32,ne=[24576,16,1,1],k=16,ties=0 60.23 46.36 1.299x 23.03%
12 type=f32,ne=[32768,1,1,1],k=16,ties=0 45.50 15.37 2.960x 66.22%
13 type=f32,ne=[32768,16,1,1],k=16,ties=0 74.55 56.46 1.320x 24.27%
14 type=f32,ne=[65536,1,1,1],k=16,ties=0 60.21 17.49 3.443x 70.95%
15 type=f32,ne=[65536,16,1,1],k=16,ties=0 121.82 103.37 1.178x 15.15%
16 type=f32,ne=[131072,1,1,1],k=16,ties=0 61.07 25.49 2.396x 58.26%
17 type=f32,ne=[131072,16,1,1],k=16,ties=0 163.35 149.91 1.090x 8.23%
18 type=f32,ne=[1,1,1,1],k=1,ties=0 3.95 2.38 1.660x 39.75%
19 type=f32,ne=[1000,1,1,1],k=1,ties=0 15.09 2.80 5.389x 81.44%
20 type=f32,ne=[65000,1,1,1],k=1,ties=0 60.35 5.57 10.835x 90.77%
21 type=f32,ne=[200000,1,1,1],k=1,ties=0 67.98 8.28 8.210x 87.82%
22 type=f32,ne=[1,16,1,1],k=1,ties=0 3.97 2.38 1.668x 40.05%
23 type=f32,ne=[1000,16,1,1],k=1,ties=0 15.64 2.83 5.527x 81.91%
24 type=f32,ne=[65000,16,1,1],k=1,ties=0 121.39 20.44 5.939x 83.16%
25 type=f32,ne=[200000,16,1,1],k=1,ties=0 197.94 51.99 3.807x 73.73%
26 type=f32,ne=[4,1,1,1],k=4,ties=0 4.53 2.76 1.641x 39.07%
27 type=f32,ne=[1000,1,1,1],k=4,ties=0 15.13 5.04 3.002x 66.69%
28 type=f32,ne=[65000,1,1,1],k=4,ties=0 60.47 14.87 4.067x 75.41%
29 type=f32,ne=[200000,1,1,1],k=4,ties=0 72.86 27.06 2.693x 62.86%
30 type=f32,ne=[4,16,1,1],k=4,ties=0 4.56 2.75 1.658x 39.69%
31 type=f32,ne=[1000,16,1,1],k=4,ties=0 15.70 5.07 3.097x 67.71%
32 type=f32,ne=[65000,16,1,1],k=4,ties=0 121.45 95.50 1.272x 21.37%
33 type=f32,ne=[200000,16,1,1],k=4,ties=0 202.94 180.27 1.126x 11.17%
34 type=f32,ne=[8,1,1,1],k=8,ties=0 5.08 2.77 1.834x 45.47%
35 type=f32,ne=[1000,1,1,1],k=8,ties=0 15.06 5.06 2.976x 66.40%
36 type=f32,ne=[65000,1,1,1],k=8,ties=0 60.57 17.09 3.544x 71.78%
37 type=f32,ne=[200000,1,1,1],k=8,ties=0 72.95 28.22 2.585x 61.32%
38 type=f32,ne=[8,16,1,1],k=8,ties=0 5.17 2.77 1.866x 46.42%
39 type=f32,ne=[1000,16,1,1],k=8,ties=0 15.67 5.09 3.079x 67.52%
40 type=f32,ne=[65000,16,1,1],k=8,ties=0 121.64 100.41 1.211x 17.45%
41 type=f32,ne=[200000,16,1,1],k=8,ties=0 203.15 180.38 1.126x 11.21%
42 type=f32,ne=[10,1,1,1],k=10,ties=0 5.79 2.77 2.090x 52.16%
43 type=f32,ne=[1000,1,1,1],k=10,ties=0 15.09 5.06 2.982x 66.47%
44 type=f32,ne=[65000,1,1,1],k=10,ties=0 60.62 17.25 3.514x 71.54%
45 type=f32,ne=[200000,1,1,1],k=10,ties=0 69.67 29.35 2.374x 57.87%
46 type=f32,ne=[10,16,1,1],k=10,ties=0 5.94 2.77 2.144x 53.37%
47 type=f32,ne=[1000,16,1,1],k=10,ties=0 15.67 5.09 3.079x 67.52%
48 type=f32,ne=[65000,16,1,1],k=10,ties=0 121.83 100.74 1.209x 17.31%
49 type=f32,ne=[200000,16,1,1],k=10,ties=0 199.73 178.45 1.119x 10.65%
50 type=f32,ne=[16,1,1,1],k=16,ties=0 5.92 2.78 2.129x 53.04%
51 type=f32,ne=[1000,1,1,1],k=16,ties=0 15.06 5.07 2.970x 66.33%
52 type=f32,ne=[65000,1,1,1],k=16,ties=0 60.69 17.43 3.482x 71.28%
53 type=f32,ne=[200000,1,1,1],k=16,ties=0 72.89 30.63 2.380x 57.98%
54 type=f32,ne=[16,16,1,1],k=16,ties=0 5.99 2.78 2.155x 53.59%
55 type=f32,ne=[1000,16,1,1],k=16,ties=0 15.69 5.10 3.076x 67.50%
56 type=f32,ne=[65000,16,1,1],k=16,ties=0 121.86 102.60 1.188x 15.81%
57 type=f32,ne=[200000,16,1,1],k=16,ties=0 203.27 181.80 1.118x 10.56%
58 type=f32,ne=[32,1,1,1],k=32,ties=0 6.93 2.79 2.484x 59.74%
59 type=f32,ne=[1000,1,1,1],k=32,ties=0 15.11 5.09 2.969x 66.31%
60 type=f32,ne=[65000,1,1,1],k=32,ties=0 61.22 19.82 3.089x 67.62%
61 type=f32,ne=[200000,1,1,1],k=32,ties=0 73.10 35.55 2.056x 51.37%
62 type=f32,ne=[32,16,1,1],k=32,ties=0 7.03 2.79 2.520x 60.31%
63 type=f32,ne=[1000,16,1,1],k=32,ties=0 15.69 5.13 3.058x 67.30%
64 type=f32,ne=[65000,16,1,1],k=32,ties=0 122.44 112.04 1.093x 8.49%
65 type=f32,ne=[200000,16,1,1],k=32,ties=0 203.47 183.00 1.112x 10.06%
66 type=f32,ne=[40,1,1,1],k=40,ties=0 7.86 2.80 2.807x 64.38%
67 type=f32,ne=[1000,1,1,1],k=40,ties=0 15.12 4.61 3.280x 69.51%
68 type=f32,ne=[65000,1,1,1],k=40,ties=0 61.58 22.16 2.779x 64.01%
69 type=f32,ne=[200000,1,1,1],k=40,ties=0 73.27 74.41 0.985x -1.56%
70 type=f32,ne=[40,16,1,1],k=40,ties=0 8.17 2.81 2.907x 65.61%
71 type=f32,ne=[1000,16,1,1],k=40,ties=0 15.62 4.65 3.359x 70.23%
72 type=f32,ne=[65000,16,1,1],k=40,ties=0 122.72 112.35 1.092x 8.45%
73 type=f32,ne=[200000,16,1,1],k=40,ties=0 203.42 181.39 1.121x 10.83%
74 type=f32,ne=[400,1,1,1],k=400,ties=0 11.95 3.33 3.589x 72.13%
75 type=f32,ne=[1000,1,1,1],k=400,ties=0 15.13 5.12 2.955x 66.16%
76 type=f32,ne=[65000,1,1,1],k=400,ties=0 64.54 46.44 1.390x 28.04%
77 type=f32,ne=[200000,1,1,1],k=400,ties=0 69.35 70.52 0.983x -1.69%
78 type=f32,ne=[400,16,1,1],k=400,ties=0 12.35 3.35 3.687x 72.87%
79 type=f32,ne=[1000,16,1,1],k=400,ties=0 15.71 5.16 3.045x 67.15%
80 type=f32,ne=[65000,16,1,1],k=400,ties=0 126.35 115.77 1.091x 8.37%
81 type=f32,ne=[200000,16,1,1],k=400,ties=0 200.56 177.46 1.130x 11.52%

Requirements

@pwilkin
pwilkin requested review from a team and IMbackK as code owners September 3, 2026 10:52
@github-actions github-actions Bot added ggml changes relating to the ggml tensor library for machine learning CUDA Related to the CUDA backend labels Sep 3, 2026
@pwilkin

pwilkin commented Sep 3, 2026

Copy link
Copy Markdown
Member Author

Supersedes #27466 and #26592

@IMbackK

IMbackK commented Sep 3, 2026 •

Copy link
Copy Markdown
Contributor

Using wave64 mode in rdna gpus is not supported from hip. I know its tempting, the gpus can dual issue valu instructions in this mode without requireing the Compiler to find opertunities for Packed fp32 math.

See
ROCm/legacy-rocm-build#4121 (comment)

There are various mines to step on.

@pwilkin

pwilkin commented Sep 3, 2026 •

Copy link
Copy Markdown
Member Author

@IMbackK I know. But the issue is, HIP is behind due to this. So the question is - do we take the performance loss by following the guidelines or do we just risk it and optimize for actual architectures?

The more I'm looking into the drivers, the more I'm convinced that AMD is shooting themselves in the foot like this. The Vulkan backend is using wavefront64 and it gives them numerous decode advantages already.

@pwilkin

pwilkin commented Sep 3, 2026

Copy link
Copy Markdown
Member Author

But of course if you say we're not using wave64, I'll revert to the "safe" wave32 kernels.

@IMbackK

IMbackK commented Sep 3, 2026 •

Copy link
Copy Markdown
Contributor

I dont know wtf they are doing either, but in this case we have to do what they say as i have tried this for work and it was an absolute maintenance nightmare over time with: header implemented cuda compat intrinsics not working correctly for the extra threads, random hip headers having a preprocessor guard preventing wave64 with gfx10+ etc.

I get that if its just one file (top_k) its probably not going to hurt too much, but using it helps in alot of places for gfx11 (indeed see ggml-vulkan or mesa/aco generally it uses wave64 almost exclusively) so i dont think we can take the precedence.

@pwilkin

pwilkin commented Sep 3, 2026

Copy link
Copy Markdown
Member Author

Aight, reverting to wave32 then. Argh.

Assisted-by: OpenAI Codex
@Geramy

Geramy commented Sep 3, 2026

Copy link
Copy Markdown
Contributor

Using wave64 mode in rdna gpus is not supported from hip. I know its tempting, the gpus can dual issue valu instructions in this mode without requireing the Compiler to find opertunities for Packed fp32 math.

See ROCm/legacy-rocm-build#4121 (comment)

There are various mines to step on.

RDNA4 supports wave64, but it operates differently on the hardware and I almost believe you loose performance from what I remember.

@0cc4m

0cc4m commented Sep 4, 2026

Copy link
Copy Markdown
Contributor

RDNA4 supports wave64, but it operates differently on the hardware and I almost believe you loose performance from what I remember.

From the Vulkan experience, we'ved switched RDNA1/2 to wave32 in some cases, but not RDNA3/4, there wave64 seemed advantageous for most kernels.

@pwilkin

pwilkin commented Sep 7, 2026

Copy link
Copy Markdown
Member Author

@IMbackK any chance you could look at it on your end? Would be nice to have the TOP_K issue resolved :)

@IMbackK

IMbackK commented Sep 7, 2026 •

Copy link
Copy Markdown
Contributor

I have pretty limited time to work on llamacpp but its on my todo list, i will get around to it.

  1. Could you rebench gfx11 in wave32 mode?
  2. i have not looked at the full code but 93ceb53 looks like its fixing wave size to 32 which is wrong for non-rdna hip devices (gcn/ cdna) if they are takeing that path.

@IMbackK

IMbackK commented Sep 15, 2026

Copy link
Copy Markdown
Contributor

@pwilkin are you still perusing getting this merged?

@pwilkin

pwilkin commented Sep 16, 2026

Copy link
Copy Markdown
Member Author

Yeah, will check your remarks.

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

CUDA Related to the CUDA backend ggml changes relating to the ggml tensor library for machine learning

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants