diff --git a/test/test_linearizer.py b/test/test_linearizer.py index f6c3c446bc..312fdba932 100644 --- a/test/test_linearizer.py +++ b/test/test_linearizer.py @@ -11,7 +11,7 @@ from tinygrad.shape.view import View from tinygrad.tensor import Tensor, _to_np_dtype from tinygrad.engine.realize import run_schedule, lower_schedule, CompiledRunner, get_program from tinygrad.codegen.opt.heuristic import hand_coded_optimizations -from tinygrad.helpers import prod, Context, getenv, CI, flatten, dedup, AMX, AMD_LLVM +from tinygrad.helpers import prod, Context, getenv, CI, flatten, dedup, AMX, AMD_LLVM, TC_SELECT, TC_OPT from tinygrad.dtype import DType, dtypes, PtrDType, AddrSpace from tinygrad.codegen import apply_rewrites, rewrites_for_views @@ -380,6 +380,7 @@ class TestLinearizer(unittest.TestCase): def test_tensor_cores_multi_reduce(self): for tc in Device[Device.DEFAULT].renderer.tensor_cores: if not is_dtype_supported(tc.dtype_in) or not is_dtype_supported(tc.dtype_out): continue + if tc.dtype_in is dtypes.bfloat16: continue # <-- broken with numpy # this will be a M=G16, N=G32, M=G16, M=G16, K=R16, K=R16, K=R16 with 9 choices of TC MNK axes golden_result = None for axis in range(9): @@ -1054,10 +1055,8 @@ def _helper_linearizer_opt_ast(realized_ast:UOp, real_bufs:list[Buffer], opts=[] def check_opt(opts, create_k, expected_color_size): k = create_k() lins.append(k) - if apply_tc: - assert k.apply_tensor_cores(1, extra_opts=opts), "no tensor core triggered" - else: - k.apply_opts(opts) + if apply_tc: k.apply_opt(Opt(OptOps.TC, 0, (TC_SELECT.value, TC_OPT.value, 1))) + k.apply_opts(opts) if expected_color_size is not None: cs = list(zip(k.colors(), k.full_shape)) assert cs == expected_color_size, f"expected={expected_color_size} got={cs}" @@ -1193,24 +1192,6 @@ class TestKernelOpts(unittest.TestCase): Opt(OptOps.UPCAST, 0, 2)], # No globals ]) - @unittest.skipUnless(Device[Device.DEFAULT].renderer.tensor_cores, "test requires tensor cores") - @unittest.skipUnless(Device[Device.DEFAULT].renderer.has_local, "test requires locals") - def test_invalid_tensor_core_extra_opts(self): - N = 128 - Tensor.manual_seed(1552) - a = Tensor.rand(N, N) - b = Tensor.rand(N, N) - realized_ast, _ = helper_realized_ast(a@b) - invalid_opts = [ - [Opt(OptOps.LOCAL, 2, 2)], - [Opt(OptOps.UPCAST, 2, 2)], - [Opt(OptOps.LOCAL, 0, 2), Opt(OptOps.LOCAL, 2, 2)], - ] - for x in invalid_opts: - k = Kernel(realized_ast) - with self.assertRaises(AssertionError): - assert k.apply_tensor_cores(use_tensor_cores=1, extra_opts=x), "no valid tensor core" # for METAL in runners - @unittest.skipUnless(Device[Device.DEFAULT].renderer.tensor_cores, "test requires tensor cores") @unittest.skipUnless(any(tc.dtype_in == tc.dtype_out == dtypes.half for tc in Device[Device.DEFAULT].renderer.tensor_cores), "test requires tensor cores with accumulation in half") # testing with half suffices. diff --git a/tinygrad/codegen/opt/kernel.py b/tinygrad/codegen/opt/kernel.py index 2d3af8d3d1..7efe47e28a 100644 --- a/tinygrad/codegen/opt/kernel.py +++ b/tinygrad/codegen/opt/kernel.py @@ -399,7 +399,7 @@ class Kernel: return True return False - def apply_tensor_cores(self, use_tensor_cores=1, extra_opts:list[Opt]|None=None, axis:int=0, tc_select:int|None=None, tc_opt:int|None=None) -> bool: + def apply_tensor_cores(self, use_tensor_cores=1) -> bool: # , extra_opts:list[Opt]|None=None) -> bool: """ Attempts to apply a tensor core optimization to the kernel. If one exists and applies properly, return true, otherwise return false. Tensor cores are optimized instructions that matrix multiply-accumulate across a wave of threads: D(M, N) = A(M, K) * B(K, N) + C(M, N). @@ -417,23 +417,19 @@ class Kernel: 1: allows kernels with multiple reduce axes and also multiplication of Ops.CAST'd buffers 2: allows kernels with M, N, K axes that are not multiples of the tensor core dimensions by applying padding those axes as needed """ - if tc_select is None: tc_select = TC_SELECT.value - if tc_opt is None: tc_opt = TC_OPT.value if not self.opts.tensor_cores: return False try: # check TC first and apply hand-coded opts if successful - self.apply_opt(Opt(OptOps.TC, axis, (tc_select, tc_opt, use_tensor_cores))) + self.apply_opt(Opt(OptOps.TC, 0, (TC_SELECT.value, TC_OPT.value, use_tensor_cores))) if (tc_opts:=self.tensor_core_opts) is not None: - if extra_opts is not None: self.apply_opts(extra_opts) - else: - if AMX: return True # skip hand-coded TC opts if AMX, upcasting will make kernel slower - # hand-coded TC opts - for tc_dim in [tc_dim for tc_dim in [1,0] if tc_opts.axes_exist[tc_dim]]: # attempt to upcast M and N - szs = [sz for sz in [5,4,3,2] if self.full_shape[tc_opts.axes[tc_dim]] % sz == 0] - if szs: self.apply_opt(Opt(OptOps.UPCAST, tc_opts.axes[tc_dim], szs[0])) + if AMX: return True # skip hand-coded TC opts if AMX, upcasting will make kernel slower + # hand-coded TC opts + for tc_dim in [tc_dim for tc_dim in [1,0] if tc_opts.axes_exist[tc_dim]]: # attempt to upcast M and N + szs = [sz for sz in [5,4,3,2] if self.full_shape[tc_opts.axes[tc_dim]] % sz == 0] + if szs: self.apply_opt(Opt(OptOps.UPCAST, tc_opts.axes[tc_dim], szs[0])) - if tc_opts.axes_exist[0] and (szs := [sz for sz in [4,2] if self.full_shape[tc_opts.axes[0]] % sz == 0]): # attempt to local N - self.apply_opt(Opt(OptOps.LOCAL, tc_opts.axes[0], szs[0])) + if tc_opts.axes_exist[0] and (szs := [sz for sz in [4,2] if self.full_shape[tc_opts.axes[0]] % sz == 0]): # attempt to local N + self.apply_opt(Opt(OptOps.LOCAL, tc_opts.axes[0], szs[0])) return True except KernelOptError: return False