diff --git a/test/test_linearizer.py b/test/test_linearizer.py index ef1f6150eb..814274b73d 100644 --- a/test/test_linearizer.py +++ b/test/test_linearizer.py @@ -749,7 +749,7 @@ class TestLinearizer(unittest.TestCase): k = helper_linearizer_opt(r, [[Opt(OptOps.UNROLL, 0, 4)]], apply_tc=True, atol=3e-2, rtol=1e-3)[-1] for u in k.uops: if u.op is UOps.WMMA: - assert u.src[-1].dtype == dtypes.float.vec(prod(tc.thread_local_sizes[2])) + #assert u.src[-1].dtype == dtypes.float.vec(prod(tc.thread_local_sizes[2])) assert u.src[-1].src[0].op != UOps.PHI @unittest.skipUnless(Device[Device.DEFAULT].renderer.tensor_cores, "test requires tensor cores") @@ -761,7 +761,7 @@ class TestLinearizer(unittest.TestCase): k = helper_linearizer_opt(r, [[Opt(OptOps.UNROLL, 0, 4)]], apply_tc=True, atol=3e-2, rtol=1e-3)[-1] for u in k.uops: if u.op is UOps.WMMA: - assert u.src[-1].dtype == dtypes.float.vec(prod(tc.thread_local_sizes[2])) + #assert u.src[-1].dtype == dtypes.float.vec(prod(tc.thread_local_sizes[2])) assert u.src[-1].src[0].op != UOps.PHI @unittest.skipUnless(Device[Device.DEFAULT].renderer.supports_float4, "test requires float4") diff --git a/tinygrad/codegen/kernel.py b/tinygrad/codegen/kernel.py index 463e16aef2..6b18eb4e84 100644 --- a/tinygrad/codegen/kernel.py +++ b/tinygrad/codegen/kernel.py @@ -146,6 +146,9 @@ class Kernel: def first_reduce(self) -> int: return [x!=y for x,y in zip(self.sts[0].shape[:self.shape_len-self.upcasted]+(0,), self.full_shape[:self.shape_len-self.upcasted]+(1,))].index(True) # noqa: E501 + @property + def first_upcast(self) -> int: return self.shape_len-self.upcasted + @property def reduceop(self) -> Optional[LazyOp]: return self.reduceops[0] if len(self.reduceops) > 0 else None @@ -666,21 +669,16 @@ class Kernel: return st1.reshape(new_shape).simplify().permute(tuple(permaxis)).reshape(st1.shape).simplify() if self.opts.device in {"AMD", "HIP"}: - reduce_axes = [self.shape_len-self.upcasted] - upcast_axis: Tuple[Tuple[Tuple[int, int], ...], Tuple[Tuple[int, int], ...], Tuple[Tuple[int, int], ...]] = \ - (((self.shape_len-self.upcasted, 16),), ((self.shape_len-self.upcasted, 16),), ((self.shape_len-self.upcasted+1, 8),)) + reduce_axes, upcast_axis = [0], [[(0, 16)], [(0, 16)], [(1, 8)]] # https://gpuopen.com/learn/wmma_on_rdna3/ fix_st1 = functools.partial(fix_st, (8,2,2), (16,8), (16,2,4), ((1,2), (0,2), (1,1), (0,1)), ((1,0), (0,0))) fix_st2 = None elif self.opts.device == "METAL": - reduce_axes = [self.shape_len-self.upcasted] - upcast_axis = (((self.shape_len-self.upcasted+1, 2),), ((self.shape_len-self.upcasted+1, 2),), ((self.shape_len-self.upcasted+1, 2),)) + reduce_axes, upcast_axis = [0], [[(1, 2)], [(1, 2)], [(1, 2)]] fix_st1 = functools.partial(fix_st, (2,4,2,2), (8,2), (2,2,2,2), ((1,1), (0,1), (1,0), (0,3)), ((0,0), (0,2), (1,3), (1,2))) fix_st2 = functools.partial(fix_st, (2,4,2,2), (8,2), (2,2,2,2), ((0,0), (1,1), (1,2), (0,2), (1,0)), ((0,1), (0,3), (1,3))) elif self.opts.device in {"CUDA", "NV"}: - reduce_axes = [self.shape_len-self.upcasted, self.shape_len-self.upcasted+1] - upcast_axis = (((self.shape_len-self.upcasted, 8),), ((self.shape_len-self.upcasted+2, 2), (self.shape_len-self.upcasted+3, 2)), - ((self.shape_len-self.upcasted+2, 2), (self.shape_len-self.upcasted+3, 2))) + reduce_axes, upcast_axis = [0, 1], [[(0, 8)], [(2, 2), (3, 2)], [(2, 2), (3, 2)]] # https://docs.nvidia.com/cuda/parallel-thread-execution/#warp-level-matrix-fragment-mma-16816-float fix_st1 = functools.partial(fix_st, (2,2,2,2,2), (8,2,2,2), (2,2,2,2,2,2), ((1,1), (1,0), (0,2), (0,3), (0,4)), ((1,3), (1,4), (1,2), (0,0), (0,1), (1,5))) @@ -690,11 +688,11 @@ class Kernel: raise RuntimeError("unsupported device for tensor cores") assert apply_to_st is None, "double tensor core? not supported" - wmma_sz = [prod(l) for l in tc.thread_local_sizes] - wmma_arg = (str(tc), tc.dims, tc.dtype_in, tc.dtype_out, tuple(wmma_sz), self.opts.device, upcast_axis, tuple(reduce_axes)) + wmma_arg = (str(tc), tc.dims, tc.dtype_in, tc.dtype_out, self.opts.device, + tuple(tuple((self.first_upcast+ax, sz) for ax, sz in up) for up in upcast_axis), + tuple(self.first_upcast+ax for ax in reduce_axes)) ret = LazyOp(ReduceOps.WMMA, (fixup_ast(rsrc.src[0], fix_st1), fixup_ast(rsrc.src[1], fix_st2)), wmma_arg) - new_reduce_axes = tuple(i for i in arg if i not in reduce_axes) - return LazyOp(op.op, (ret,), new_reduce_axes) if new_reduce_axes else ret + return LazyOp(op.op, (ret,), new_reduce_axes) if (new_reduce_axes:=tuple(i for i in arg if i-self.first_upcast not in reduce_axes)) else ret if self.group_for_reduces: start = LazyOp(op.op, tuple(fixup_ast(x, apply_to_st) for x in op.src), arg) local_shape = (1,) * self.global_dims + self.full_shape[self.global_dims:self.global_dims+self.local_dims+self.group_for_reduces] + \ diff --git a/tinygrad/codegen/lowerer.py b/tinygrad/codegen/lowerer.py index bf00a30a4d..c454af7c6d 100644 --- a/tinygrad/codegen/lowerer.py +++ b/tinygrad/codegen/lowerer.py @@ -7,7 +7,7 @@ from tinygrad.dtype import dtypes, PtrDType, ImageDType, DType from tinygrad.ops import BufferOps, LazyOp, TernaryOps, ReduceOps, UnaryOps, MetaOps, KernelInfo, MemBuffer from tinygrad.codegen.uops import UOp, UOps from tinygrad.renderer import Renderer -from tinygrad.helpers import getenv, all_int, get_contraction +from tinygrad.helpers import getenv, all_int, get_contraction, prod # TODO: this needs to be replaced, there shouldn't be variables in the shapetracker, only ints and UOps from tinygrad.shape.symbolic import Variable, NumNode, SumNode, MulNode, DivNode, ModNode, LtNode, AndNode @@ -171,7 +171,8 @@ class IndependentLowerer: if x.op in ReduceOps: dtype = x.dtype.base if isinstance(x.dtype, ImageDType) else x.dtype if x.op is ReduceOps.WMMA: - wmma_sz, upcast_axis = x.arg[4], x.arg[6] + upcast_axis = x.arg[-2] + wmma_sz = [prod(x[1] for x in l) for l in upcast_axis] ret = UOp(UOps.WMMA, dtype=dtype.vec(wmma_sz[2]), src=( UOp(UOps.CONTRACT, dtype=cast(DType, in_uops[0].dtype).vec(wmma_sz[0]), src=(in_uops[0],), arg=upcast_axis[0]), UOp(UOps.CONTRACT, dtype=cast(DType, in_uops[1].dtype).vec(wmma_sz[1]), src=(in_uops[1],), arg=upcast_axis[1]), diff --git a/tinygrad/renderer/__init__.py b/tinygrad/renderer/__init__.py index 94396c8ca6..87cf257b5a 100644 --- a/tinygrad/renderer/__init__.py +++ b/tinygrad/renderer/__init__.py @@ -12,7 +12,6 @@ class TensorCore: # D = A * B + C, A is (M x K), B is (K x N), C and D are (M x dtype_in: DType # dtype for A and B dtype_out: DType # dtype for C and D threads: List[Tuple[int,int]] # list of (TC dim,amt) that construct the warp thread structure - thread_local_sizes: List[List[int]] # in each thread, the number of elements stored in registers for each TC dim def __str__(self): return "_".join(["WMMA"] + list(map(str, self.dims)) + [self.dtype_in.name, self.dtype_out.name]) @dataclass(frozen=True) diff --git a/tinygrad/renderer/assembly.py b/tinygrad/renderer/assembly.py index c5b25ccdfd..9931e9c19b 100644 --- a/tinygrad/renderer/assembly.py +++ b/tinygrad/renderer/assembly.py @@ -21,7 +21,7 @@ class PTXRenderer(Renderer): global_max = (2147483647, 65535, 65535) local_max = (1024, 1024, 64) shared_max = 49152 - tensor_cores = [TensorCore(dims=(8,16,16), threads=[(0,2),(0,2),(1,2),(1,2),(0,2)], thread_local_sizes=[[2,2,2],[2,2],[2,2]], dtype_in=di, dtype_out=do) for (di, do) in ([(dtypes.half, dtypes.float)])] # noqa: E501 + tensor_cores = [TensorCore(dims=(8,16,16), threads=[(0,2),(0,2),(1,2),(1,2),(0,2)], dtype_in=di, dtype_out=do) for (di, do) in ([(dtypes.half, dtypes.float)])] # noqa: E501 def __init__(self, arch:str, device="CUDA"): self.device, self.tensor_cores = device, PTXRenderer.tensor_cores if int(arch[3:]) >= 80 else [] # language options diff --git a/tinygrad/renderer/cstyle.py b/tinygrad/renderer/cstyle.py index 5dd07ea7cd..ad36e2951b 100644 --- a/tinygrad/renderer/cstyle.py +++ b/tinygrad/renderer/cstyle.py @@ -215,7 +215,7 @@ class OpenCLRenderer(CStyleLanguage): class MetalRenderer(CStyleLanguage): device = "METAL" shared_max = 32768 - tensor_cores = [TensorCore(dims=(8,8,8), threads=[(0,2),(1,4),(0,2),(1,2)], thread_local_sizes=[[2],[2],[2]], dtype_in=di, dtype_out=do) for (di, do) in [(dtypes.float, dtypes.float), (dtypes.half, dtypes.float), (dtypes.half, dtypes.half)]] # noqa: E501 + tensor_cores = [TensorCore(dims=(8,8,8), threads=[(0,2),(1,4),(0,2),(1,2)], dtype_in=di, dtype_out=do) for (di, do) in [(dtypes.float, dtypes.float), (dtypes.half, dtypes.float), (dtypes.half, dtypes.half)]] # noqa: E501 def __init__(self): self.tensor_cores = MetalRenderer.tensor_cores if os.uname().machine == "arm64" else [] # language options @@ -265,7 +265,7 @@ class CUDARenderer(CStyleLanguage): global_max = (2147483647, 65535, 65535) local_max = (1024, 1024, 64) shared_max = 49152 - tensor_cores = [TensorCore(dims=(8,16,16), threads=[(0,2),(0,2),(1,2),(1,2),(0,2)], thread_local_sizes=[[2,2,2],[2,2],[2,2]], dtype_in=di, dtype_out=do) for (di, do) in ([(dtypes.half, dtypes.float), (dtypes.bfloat16, dtypes.float)])] # noqa: E501 + tensor_cores = [TensorCore(dims=(8,16,16), threads=[(0,2),(0,2),(1,2),(1,2),(0,2)], dtype_in=di, dtype_out=do) for (di, do) in ([(dtypes.half, dtypes.float), (dtypes.bfloat16, dtypes.float)])] # noqa: E501 def __init__(self, arch:str): self.tensor_cores = CUDARenderer.tensor_cores if int(arch[3:]) >= 80 else [] # language options @@ -325,7 +325,7 @@ def _make_hip_dtype(base_type, name, cnt): class AMDRenderer(CStyleLanguage): device = "AMD" shared_max = 65536 - tensor_cores = [TensorCore(dims=(16,16,16), threads=[(0,8),(0,2),(1,2)], thread_local_sizes=[[16],[16],[4,2]], dtype_in=di, dtype_out=do) for (di, do) in [(dtypes.half, dtypes.float), (dtypes.half, dtypes.half)]] # noqa: E501 + tensor_cores = [TensorCore(dims=(16,16,16), threads=[(0,8),(0,2),(1,2)], dtype_in=di, dtype_out=do) for (di, do) in [(dtypes.half, dtypes.float), (dtypes.half, dtypes.half)]] # noqa: E501 # language options kernel_prefix = """extern "C" __attribute__((device)) __attribute__((const)) size_t __ockl_get_local_id(unsigned int); diff --git a/tinygrad/runtime/ops_python.py b/tinygrad/runtime/ops_python.py index f523a3ae6c..4c1a3e3683 100644 --- a/tinygrad/runtime/ops_python.py +++ b/tinygrad/runtime/ops_python.py @@ -150,13 +150,13 @@ class PythonProgram: return out # TODO: refactor these to a shared TensorCoreLayout in kernel.py - if arg[5] == "METAL": + if arg[4] == "METAL": # A (2 elements on 32 threads): row major def a_b_elem(x, i, j, goff): return x[(i%2)][goff+(i//2)%2+(j%4)*2+(i//4)*8+(j//4)*16] # (i, j), C, D (2 elements on 32 threads): row major same as A/B def c_map(lane, elem): return (elem + ((lane%2)*2) + ((lane//8)%2)*4, ((lane//2)%4) + (lane//16)*4) ul[i] = wmma_helper(32, 8, 2, 2, 2, a_b_elem, a_b_elem, c_map) - elif arg[5] == "AMD": + elif arg[4] == "AMD": # A (16 elements on 32 threads): col major, lane 16-32 == lane 0-15 def a_elem(x, i, j, goff): assert x[i][goff+j] == x[i][goff+j+16], "warp elements not duplicated properly across lanes" @@ -165,7 +165,7 @@ class PythonProgram: def b_elem(x, i, j, goff): return a_elem(x, j, i, goff) # pylint: disable=arguments-out-of-order def c_map(lane, elem): return (lane%16, lane//16+elem*2) # (i, j), C, D (8 elements on 32 threads): row major ul[i] = wmma_helper(32, 16, 16, 16, 8, a_elem, b_elem, c_map) - elif arg[5] == "CUDA": + elif arg[4] == "CUDA": # A (8 elements on 32 threads) def a_elem(x, i, j, goff): return x[(i%2)+(j//8)*2+(i//8)*4][goff+((i//2)%4)+(j%8)*4] # B (4 elements on 32 threads)