diff --git a/extra/gemm/metal_matmul.py b/extra/gemm/metal_matmul.py index c9782bffba..7c5f7774a8 100644 --- a/extra/gemm/metal_matmul.py +++ b/extra/gemm/metal_matmul.py @@ -2,7 +2,9 @@ import numpy as np from tinygrad.helpers import dtypes from tinygrad.runtime.ops_metal import RawMetalBuffer, MetalProgram + N = 2048 + a = RawMetalBuffer(N*N, dtypes.float32) nb = np.random.default_rng().standard_normal(size=(N,N), dtype=np.float32) #.astype(np.int32).astype(np.float32) @@ -81,7 +83,7 @@ kernel void test(device float *a, device const float *data1, device const float }} }} }}""") -tm = min([prog([N*N//(2*4*4)], [4*32], a, b, c, wait=True) for _ in range(10)]) +tm = min([prog([N*N//(2*4*4)], [4*32], a, b, c, wait=True) for _ in range(20)]) na = a.toCPU().reshape(N,N) comp = nb@nc if N <= 32: @@ -90,3 +92,14 @@ if N <= 32: print(f"{N*N:10d} {tm*1e6:9.2f} us, would be {FLOPS*1e-9/tm:.2f} GFLOPS matmul") np.testing.assert_allclose(na, comp, atol=1e-3) +import time, torch +b = torch.from_numpy(nb).to('mps') +c = torch.from_numpy(nc).to('mps') + +def torch_prog(b, c): + st = time.perf_counter() + a = b@c + a.cpu() + return time.perf_counter() - st +tm = min([torch_prog(b, c) for _ in range(20)]) +print(f"{N*N:10d} {tm*1e6:9.2f} us, would be {FLOPS*1e-9/tm:.2f} GFLOPS matmul in torch")