[BUG]: select_config estimates matmul wave counts with A100's 108 SMs on every GPU; the target's SM count is 7-10 % faster (median) on A10 and L4
I maintainer di solito rispondono entro 1 giorno
Nessuno ha ancora preso questa issue.
Valutazione
- Difficoltà
- 2/5
- Tempo stimato
- 1-3 ore
- Idoneità per principianti
- 78/100
- Tipo di issue
- Bug
- Chiarezza
- Specificata chiaramente
- Stato di attività
- Attiva
- Ambito
- machine-learning, performance
Direzione di ricerca
Inizia in max/kernels/src/linalg/utils_gpu.mojo dai due calcoli del numero di wave di select_config intorno alle righe 488 e 521, quindi esamina gpu_info e _shared_memory_usage nelle vicinanze. Riproduci il caso A10 con una matmul bf16 con M=256, N=K=4096 e LOGGING_LEVEL=INFO. Il lavoro è completato quando l’euristica usa il numero di SM della GPU target e la configurazione registrata corrisponde al target reale anziché ai 108 SM di A100.
Scritto dal modello di indicizzazione a partire dal testo della issue.
Descrizione
Bug description
select_config picks the block tile and the number of split-K partitions for every 16-bit matmul that has no entry in the Ampere tuning table (on NVIDIA GPUs other than A100: all of them) by estimating how many "waves" of blocks it takes to cover the SMs, and it does that with A100.sm_count (108) whatever the GPU is:
The comment right above admits it ("We use A100 properties for target="gpu" (default) on A10, L4"). A10 has 72 SMs and L4 58 (_builtin_targets.mojo), so the wave count is wrong on both, and with it the split-K decision: replaying the function with the real count, 25 of the 40 (M, N, K) points below change their config on A10 and 23 on L4, almost always the partition count, in both directions.
What I measured, bf16, transpose_b=True, bench_matmul.mojo's kernel with kernels built from source at f5eda372 (the two sites are unchanged on main), Llama-3-8B shapes (N×K = 4096×4096, 6144×4096, 4096×14336, 14336×4096, 28672×4096) at M = 16 … 2048, native Linux. Same tree for both arms, with the divisor selected at compile time (108 = today's behaviour; 72 / 58 = the real count), binaries interleaved per shape, two rounds, minimum per iteration over ~300 iterations; the 15–17 shapes whose config does not change are an A/A control inside the run (median 1.00× on both GPUs), and an A100 run with 108 in both arms gives 1.000×. The configs the GPU actually dispatched were read back with -D LOGGING_LEVEL=INFO and match the replay 16/16.
Speedup = t(108) / t(real count), on the shapes whose config changes:
| A10 (72 SMs), run 1 | A10, independent run 2 | L4 (58 SMs) | |
|---|---|---|---|
| shapes that change config | 25 / 40 | 25 / 40 | 23 / 40 |
| median speedup | 1.074× | 1.071× | 1.096× |
| range | 0.967 – 1.362× | 0.972 – 1.358× | 0.773 – 1.710× |
| ≥ 1.05× / ≤ 0.95× | 16 / 0 | 16 / 0 | 15 / 1 |
| control shapes (unchanged config), median | 0.999× | 0.999× | 1.000× |
The largest wins are where 108 under-counts the waves and over-splits: 4096×4096 with M = 256–1024 on A10 (3 → 1 partitions, 1.18–1.36×) and the M ≤ 128 rows of 4096×4096 / 6144×4096 on L4 (1.29–1.71×). The one regression is (1024, 4096, 14336) on L4: 2 → 1 partitions, 0.94× with the mean per iteration (0.77× with the minimum). There the num_waves_base > 3 cutoff at L500, now counted in real waves (256 blocks / 58 SMs), switches split-K off where it still pays; that threshold was tuned for 108 too, and I left it alone.
All 28 shapes whose config changes on either GPU (tile is BM×BN, pN = split-K partitions; A10 speedup as run 1 / run 2, L4 as min / mean)
| shape (M, N, K) | A10: config 108 → 72 | A10 t108 → t72 ms | A10 speedup | L4: config 108 → 58 | L4 t108 → t58 ms | L4 speedup |
|---|---|---|---|---|---|---|
| (16, 4096, 4096) | same | 64x256 p4 → 64x256 p3 | 0.043 → 0.030 | 1.45× / 1.39× | ||
| (32, 4096, 4096) | same | 64x256 p4 → 64x256 p3 | 0.044 → 0.031 | 1.44× / 1.40× | ||
| (64, 4096, 4096) | same | 64x256 p4 → 64x256 p3 | 0.047 → 0.033 | 1.44× / 1.39× | ||
| (128, 4096, 4096) | 128x128 p3 → 128x128 p2 | 0.106 → 0.094 | 1.13× / 1.16× | same | ||
| (256, 4096, 4096) | 128x128 p3 → 128x128 p1 | 0.131 → 0.096 | 1.36× / 1.36× | same | ||
| (512, 4096, 4096) | 128x128 p3 → 128x128 p1 | 0.229 → 0.180 | 1.27× / 1.25× | 128x128 p3 → 128x128 p2 | 0.231 → 0.219 | 1.06× / 1.08× |
| (1024, 4096, 4096) | 128x128 p2 → 128x128 p1 | 0.431 → 0.364 | 1.18× / 1.17× | 128x128 p2 → 128x128 p1 | 0.448 → 0.406 | 1.10× / 1.15× |
| (16, 4096, 14336) | 64x256 p6 → 64x256 p4 | 0.276 → 0.270 | 1.02× / 1.02× | 64x256 p6 → 64x256 p7 | 0.525 → 0.525 | 1.00× / 1.00× |
| (32, 4096, 14336) | 64x256 p6 → 64x256 p4 | 0.281 → 0.273 | 1.03× / 1.03× | 64x256 p6 → 64x256 p7 | 0.533 → 0.534 | 1.00× / 1.00× |
| (64, 4096, 14336) | 64x256 p6 → 64x256 p4 | 0.291 → 0.280 | 1.04× / 1.04× | 64x256 p6 → 64x256 p7 | 0.551 → 0.554 | 0.99× / 0.99× |
| (128, 4096, 14336) | 128x128 p3 → 128x128 p2 | 0.312 → 0.286 | 1.09× / 1.10× | same | ||
| (256, 4096, 14336) | 128x128 p3 → 128x128 p1 | 0.364 → 0.320 | 1.14× / 1.14× | same | ||
| (512, 4096, 14336) | 128x128 p3 → 128x128 p1 | 0.692 → 0.639 | 1.08× / 1.08× | 128x128 p3 → 128x128 p2 | 0.828 → 0.838 | 0.99× / 1.01× |
| (1024, 4096, 14336) | 128x128 p2 → 128x128 p1 | 1.377 → 1.294 | 1.06× / 1.07× | 128x128 p2 → 128x128 p1 | 1.533 → 1.984 | 0.77× / 0.94× |
| (16, 6144, 4096) | 64x256 p4 → 64x256 p3 | 0.125 → 0.123 | 1.02× / 1.02× | 64x256 p4 → 64x256 p2 | 0.087 → 0.060 | 1.44× / 1.36× |
| (32, 6144, 4096) | 64x256 p4 → 64x256 p3 | 0.129 → 0.126 | 1.02× / 1.02× | 64x256 p4 → 64x256 p2 | 0.144 → 0.103 | 1.40× / 1.29× |
| (64, 6144, 4096) | 64x256 p4 → 64x256 p3 | 0.137 → 0.131 | 1.05× / 1.04× | 64x256 p4 → 64x256 p2 | 0.228 → 0.177 | 1.29× / 1.27× |
| (128, 6144, 4096) | 128x128 p2 → 128x128 p3 | 0.148 → 0.149 | 0.99× / 1.00× | 128x128 p2 → 128x128 p1 | 0.242 → 0.141 | 1.71× / 1.58× |
| (256, 6144, 4096) | 128x128 p1 → 128x128 p3 | 0.182 → 0.188 | 0.97× / 0.97× | same | ||
| (16, 14336, 4096) | 64x256 p3 → 64x256 p1 | 0.279 → 0.264 | 1.05× / 1.05× | 64x256 p3 → 64x256 p1 | 0.528 → 0.512 | 1.03× / 1.03× |
| (32, 14336, 4096) | 64x256 p3 → 64x256 p1 | 0.287 → 0.266 | 1.08× / 1.08× | 64x256 p3 → 64x256 p1 | 0.542 → 0.516 | 1.05× / 1.05× |
| (64, 14336, 4096) | 64x256 p3 → 64x256 p1 | 0.304 → 0.269 | 1.13× / 1.12× | 64x256 p3 → 64x256 p1 | 0.564 → 0.522 | 1.08× / 1.08× |
| (128, 14336, 4096) | 128x128 p3 → 128x128 p1 | 0.338 → 0.288 | 1.17× / 1.18× | 128x128 p3 → 128x128 p1 | 0.614 → 0.537 | 1.15× / 1.14× |
| (256, 14336, 4096) | 128x128 p2 → 128x128 p1 | 0.419 → 0.385 | 1.09× / 1.08× | 128x128 p2 → 128x128 p1 | 0.666 → 0.560 | 1.19× / 1.12× |
| (16, 28672, 4096) | 64x256 p3 → 64x256 p1 | 0.549 → 0.523 | 1.05× / 1.04× | 64x256 p3 → 64x256 p1 | 1.061 → 1.022 | 1.04× / 1.04× |
| (32, 28672, 4096) | 64x256 p3 → 64x256 p1 | 0.565 → 0.526 | 1.07× / 1.07× | 64x256 p3 → 64x256 p1 | 1.089 → 1.029 | 1.06× / 1.06× |
| (64, 28672, 4096) | 64x256 p3 → 64x256 p1 | 0.597 → 0.532 | 1.12× / 1.11× | 64x256 p3 → 64x256 p1 | 1.142 → 1.041 | 1.10× / 1.09× |
| (128, 28672, 4096) | 128x128 p2 → 128x128 p1 | 0.630 → 0.595 | 1.06× / 1.06× | 128x128 p2 → 128x128 p1 | 1.187 → 1.070 | 1.11× / 1.11× |
The other 12 (M, N, K) points (M ≥ 512 on 6144×4096 / 14336×4096 / 28672×4096, M = 2048 everywhere) keep their config on both GPUs; their ratio is 0.92–1.01× on A10 with the minimum per iteration (0.996–1.011× with the mean: the A10 runs at its power limit and the minimum of the long kernels wanders) and 0.98–1.01× on L4.
Steps to reproduce
No crash to reproduce; the decision can be read back. Build any bf16 _matmul_gpu call with -D LOGGING_LEVEL=INFO and run it on an A10 with M = 256, N = K = 4096: the log shows kernel_bfloat16_bfloat16_128x128_4_NT with K partitions: 3: 64 blocks of 128×128 are one wave on 108 SMs, so three partitions still fit the two-wave budget of the heuristic. Counted on 72 SMs, three partitions would be three waves and two gain nothing, so it dispatches K partitions: 1, and the same GEMM takes 0.096 ms instead of 0.131 ms (1.36×). The Python replay of select_config I used to pick the shapes is 25 lines and matches the logged configs 16/16; I can attach it.
System information
A10 (sm_86, 72 SMs), L4 (sm_89, 58 SMs) and A100-SXM4-40GB (control) on Modal, native Linux, driver 580.95.05; Mojo 1.1.0.dev2026082405 with the kernels built from max/kernels/src at f5eda372. select_config is unchanged on main (b3e9394d).
The fix I have is two lines: gpu_info.sm_count in place of A100.sm_count at L488 and L521 (gpu_info is already ctx.default_device_info and _shared_memory_usage uses it a few lines below). That is exactly what the A10 and L4 numbers above measure, since sm_86 resolves to the A10 entry and sm_89 to L4. It is not exact for consumer parts that share a target (an RTX 3090 Ti has 84 SMs and gets the A10's 72); if you'd rather have the real count, ctx.get_attribute(DeviceAttribute.MULTIPROCESSOR_COUNT) is available to the function. I can send the PR if you'd take it.
- Lingua principale
- Mojo
- Stelle
- 29.8k
- Fork
- 3.2k
- Metriche di merge delle PR
- Nessuna PR unita negli ultimi 30g
Preparare l'ambiente
Come iniziare
- Leggi tutta la issue e poi la guida ai contributi del progetto.
- Commenta sulla issue per dire che te ne occupi tu — evita che due persone facciano lo stesso lavoro.
- Fai un fork del repository e lavora su un branch.
- Apri una pull request che faccia riferimento al numero della issue.
Altre issue di modular/modular
-
bug bug_feature_triaged max
Difficoltà 2/5 1-3 ore Idoneità per principianti 78/100
I maintainer di solito rispondono entro 1 giorno
-
[Docs] Example demonstrating `origin_of` doesn't compileForse già presa @obadafidii l’ha presa 3 giorni fa. Apertabug_feature_triaged documentation
Difficoltà 1/5 Meno di un'ora Idoneità per principianti 85/100
modular/modular#7174 · 2 commenti · 1 assegnatario ·
I maintainer di solito rispondono entro 1 giorno
-
bug_feature_triaged documentation
Difficoltà 2/5 1-3 ore Idoneità per principianti 82/100
modular/modular#7166 · 1 commento · 1 reazione ·
I maintainer di solito rispondono entro 1 giorno
-
bug_feature_triaged
Difficoltà 2/5 1-3 ore Idoneità per principianti 84/100
I maintainer di solito rispondono entro 1 giorno
-
bug_feature_triaged enhancement mojo Team: Mojo Libraries
Difficoltà 2/5 Mezza giornata Idoneità per principianti 68/100
modular/modular#7138 · 1 reazione ·
I maintainer di solito rispondono entro 1 giorno
Tutte le issue di modular/modular
Issue simili
-
Difficoltà 2/5 1-3 ore Idoneità per principianti 88/100
huggingface/diffusers#14888 ·
I maintainer di solito rispondono entro 1 giorno
-
Difficoltà 2/5 1-3 ore Idoneità per principianti 82/100
pytorch/rl#4489 · 1 commento ·
I maintainer di solito rispondono entro 1 giorno
-
bug documentation good first issue
Difficoltà 2/5 1-3 ore Idoneità per principianti 88/100
Farama-Foundation/PettingZoo#1467 · 1 commento ·
I maintainer di solito rispondono entro 1 giorno
-
Difficoltà 2/5 1-3 ore Idoneità per principianti 85/100
-
hf-audiolm-qwen: `generate_until` hardcodes `.to("cuda")` and aborts on non-CUDA acceleratorsAperta
Difficoltà 2/5 1-3 ore Idoneità per principianti 86/100
EleutherAI/lm-evaluation-harness#4256 ·
I maintainer di solito rispondono entro 1 giorno