Low GEMM Performance on Hopper GPU with Small M Shapes
Nessuno ha ancora preso questa issue.
- Lingua principale
- Cython
- Stelle
- 601
- Fork
- 46
- Metriche di merge delle PR
- Nessuna PR unita negli ultimi 30g
Descrizione
Hi,
Thank you for the great library! I’m observing some unexpected performance with GEMM on Hopper GPUs when using small M dimensions. I followed the example in example14_autotune.py.
Compared to the PyTorch implementation, the performance is significantly lower — around 30% of the expected TFLOPS and memory bandwidth utilization.
Not sure I am correctly using the API — I would greatly appreciate any suggestions.
Environment:
• GPU: H200
• CUDA: 12.8
• PyTorch: 2.6.0
• nvmath-python: 0.3.0
Benchmark code:
import argparse
import torch
import nvmath
from triton.testing import do_bench
def profile(m, n, k, dtype):
device = torch.device("cuda")
assert isinstance(device, torch.device)
X = torch.randn(m, k, device=device, dtype=dtype)
Y = torch.randn(n, k, device=device, dtype=dtype)
_torch_gemm = lambda: torch.matmul(X, Y)
mm = nvmath.linalg.advanced.Matmul(X, Y)
mm.plan(preferences={"limit":1000})
mm.autotune(iterations=1000)
# print(mm.algorithms[0].capabilities)
_nvmath_gemm = lambda: mm.execute()
t_torch = do_bench(_torch_gemm)
t_nvmath = do_bench(_nvmath_gemm)
return t_torch, t_nvmath
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="GEMM profile")
parser.add_argument("--m", type=int, default=4096)
parser.add_argument("--n", type=int, default=4096)
parser.add_argument("--k", type=int, default=4096)
args = parser.parse_args()
print("Provider,Operation,dtype,m,n,k,Runtime,GB/s,GFLOPs")
for dtype in [torch.float16, torch.bfloat16]:
t_torch, t_nvmath = profile(args.m, args.n, args.k, dtype)
m = args.m
n = args.n
k = args.k
torch_mem_bd = 2 * (m * n + n * k + m * k) * 1e3 / t_torch / 1e9
torch_gflops = 2 * m * n * k * 1e3 / t_torch / 1e9
nv_mem_bd = 2 * (m * n + n * k + m * k) * 1e3 / t_nvmath / 1e9
nv_gflops = 2 * m * n * k * 1e3 / t_nvmath / 1e9
print(f"TORCH,0,{dtype},{args.m},{args.n},{args.k},{t_torch},{torch_mem_bd},{torch_gflops}")
print(f"NVMATH,0,{dtype},{args.m},{args.n},{args.k},{t_nvmath},{nv_mem_bd},{nv_gflops}")
Guida per i contributori
Apri la guida per i contributori
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.
Direzione di ricerca
Riproduci il confronto usando examples/linalg/advanced/matmul/example14_autotune.py e il benchmark di profiling fornito nell’ambiente H200 e CUDA 12.8 indicato. Confronta la pianificazione e l’autotuning di nvmath Matmul con torch.matmul per forme M piccole; il lavoro è completato quando il divario di prestazioni è spiegato e viene documentato l’uso corretto dell’API o il fix richiesto.
Scritto dal modello di indicizzazione a partire dal testo della issue.
Valutazione
- Stack tecnologico
- python, pytorch
- Ambito
- performance
- Tipo di issue
- Bug
- Difficoltà
- 4/5
- Tempo stimato
- 3-5 giorni
- Stato di attività
- Ferma
- Chiarezza
- Abbastanza chiara
- Idoneità per principianti
- 35/100