File size: 3,337 Bytes
454b3e6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
"""Sweep tile configs for the GEMM / GEGLU / attention kernels at the shapes laya actually hits."""
import torch, sys, time, itertools, json
sys.path.insert(0, "kernels"); import tl_kernels as K
dev = "cuda"
def bench(f, n=20):
    for _ in range(3): f()
    torch.cuda.synchronize(); t = time.perf_counter()
    for _ in range(n): f()
    torch.cuda.synchronize(); return (time.perf_counter() - t) / n * 1e3
res = {"gemm": {}, "geglu": {}, "attn": {}}
Ms = [128, 1024, 8192, 28672]
cfgs = [(bm, bn, th) for bm in (64, 128) for bn in (64, 128, 256) for th in (128, 256)]
for (N, Kd) in [(2304, 768), (768, 768), (768, 1152)]:
    W = torch.randn(N, Kd, device=dev, dtype=torch.bfloat16) * 0.02; b = torch.zeros(N, device=dev)
    for M in Ms:
        A = torch.randn(M, Kd, device=dev, dtype=torch.bfloat16); C = torch.empty(M, N, device=dev, dtype=torch.bfloat16)
        tref = bench(lambda: torch.nn.functional.linear(A, W))
        best = None
        for bm, bn, th in cfgs:
            try:
                k = K.gemm_kernel(N, Kd, bm=bm, bn=bn, bk=64, stages=3, threads=th); ms = bench(lambda: k(A, W, b, C))
                if best is None or ms < best[0]: best = (ms, (bm, bn, th))
            except Exception as e: print("fail", N, Kd, M, (bm, bn, th), str(e)[:60], flush=True)
        tf = 2 * M * N * Kd / best[0] / 1e9
        print(f"gemm N={N} K={Kd} M={M}: best {best[1]} {best[0]:.3f}ms ({tf:.1f} TFLOPS)  cublas {tref:.3f}ms  (default 64x128x128thr)", flush=True)
        res["gemm"][f"{N},{Kd},{M}"] = best
for M in Ms:
    A = torch.randn(M, 768, device=dev, dtype=torch.bfloat16); Wi = torch.randn(2304, 768, device=dev, dtype=torch.bfloat16) * 0.02; C = torch.empty(M, 1152, device=dev, dtype=torch.bfloat16)
    best = None
    for bm, bn, th in cfgs:
        if bn > 128: continue
        try:
            k = K.gemm_geglu_kernel(1152, 768, bm=bm, bn=bn, bk=64, stages=3, threads=th); ms = bench(lambda: k(A, Wi, C))
            if best is None or ms < best[0]: best = (ms, (bm, bn, th))
        except Exception as e: print("fail geglu", M, (bm, bn, th), str(e)[:60], flush=True)
    print(f"geglu M={M}: best {best[1]} {best[0]:.3f}ms ({2*M*2304*768/best[0]/1e9:.1f} TFLOPS)", flush=True)
    res["geglu"][str(M)] = best
H, Dh = 12, 64
for (B, L) in [(32, 1024), (32, 128), (4, 128)]:
    qkv = torch.randn(B, L, 3, H, Dh, device=dev, dtype=torch.bfloat16); lens = torch.full((B,), L - 3, device=dev, dtype=torch.int32); O = torch.empty(B, L, H * Dh, device=dev, dtype=torch.bfloat16)
    for window in (0, 65):
        best = None
        for bm, bn, st, th in itertools.product((64, 128), (64, 128), (1, 2), (128, 256)):
            try:
                k = K.attn_kernel(B, L, H, Dh, window=window, bm=bm, bn=bn, stages=st, threads=th); ms = bench(lambda: k(qkv, lens, O))
                if best is None or ms < best[0]: best = (ms, (bm, bn, st, th))
            except Exception as e: print("fail attn", (B, L, window), (bm, bn, st, th), str(e)[:60], flush=True)
        fl = 4 * B * H * L * L * Dh if window == 0 else 4 * B * H * L * (2 * 65 + 1) * Dh
        print(f"attn B={B} L={L} window={window}: best {best[1]} {best[0]:.3f}ms ({fl/best[0]/1e9:.1f} TFLOPS)", flush=True)
        res["attn"][f"{B},{L},{window}"] = best
json.dump(res, open("kernels/tune_results.json", "w"), indent=1)
print("DONE")