Hacktoberfest 2026: le issue che i maintainer hanno segnato per ottobre, aperte e adatte ai principianti. Sfoglia le issue Hacktoberfest

[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

Aperta Adatta ai principianti
#7,180 0 commenti 0 reazioni 0 assegnatari Vedi su GitHub

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

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_feature_triaged
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:

https://github.com/modular/modular/blob/b3e9394d088cc8fc78793f663c92b381177a14ff/max/kernels/src/linalg/utils_gpu.mojo#L488

https://github.com/modular/modular/blob/b3e9394d088cc8fc78793f663c92b381177a14ff/max/kernels/src/linalg/utils_gpu.mojo#L521

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

  1. Leggi tutta la issue e poi la guida ai contributi del progetto.
  2. Commenta sulla issue per dire che te ne occupi tu — evita che due persone facciano lo stesso lavoro.
  3. Fai un fork del repository e lavora su un branch.
  4. Apri una pull request che faccia riferimento al numero della issue.

Altre issue di modular/modular

Tutte le issue di modular/modular

Issue simili

Altre issue su Machine Learning

Ricevi le nuove issue nella tua casella

Un breve riepilogo di issue GitHub adatte ai principianti.