diff --git a/tinygrad/opt/kernel.py b/tinygrad/opt/kernel.py index b39e41746f..5c51904417 100644 --- a/tinygrad/opt/kernel.py +++ b/tinygrad/opt/kernel.py @@ -1,6 +1,6 @@ from __future__ import annotations import itertools, functools, math -from dataclasses import dataclass, replace +from dataclasses import dataclass from collections import defaultdict from typing import Optional, cast, Final, Callable, Sequence from enum import Enum, auto @@ -48,7 +48,6 @@ class Kernel: def __init__(self, ast:UOp, opts:Optional[Renderer]=None): assert ast.op is Ops.SINK, ast.op self.ast = ast - self.info: KernelInfo = self.ast.arg if self.ast.arg is not None else KernelInfo() self.opts = opts if opts is not None else Device[Device.DEFAULT].renderer # verify AST matches the spec @@ -73,10 +72,11 @@ class Kernel: self.sts.append(ShapeTracker.from_shape(tuple([smax(*s) for s in zip(*[x.shape for x in self.sts])]), (0,)*self.shape_len)) # parameters for optimization - self.group_for_reduces: int = 0 self.tensor_core: Optional[TensorCore] = None self.tensor_core_opts: Optional[TensorCoreOptions] = None self.use_tensor_cores: int = 0 + self.applied_opts: list[Opt] = [] + self.dont_use_locals = False # finalized means you can't optimize anymore self.finalized: bool = False @@ -105,14 +105,12 @@ class Kernel: ret.axis_types = self.axis_types[:] # parameters for optimizations - ret.info, ret.group_for_reduces = self.info, self.group_for_reduces + ret.applied_opts, ret.dont_use_locals = self.applied_opts[:], self.dont_use_locals ret.tensor_core, ret.tensor_core_opts, ret.use_tensor_cores = self.tensor_core, self.tensor_core_opts, self.use_tensor_cores ret.finalized = self.finalized return ret - def update_info(self, **updates): self.info = replace(self.info, **updates) - @property def membufs(self) -> list[UOp]: return dedup([x.src[0].base for x in self.bufs if x.op in {Ops.LOAD, Ops.STORE}]) @@ -145,19 +143,13 @@ class Kernel: def shape_len(self) -> int: return len(self.sts[0].shape) @property - def global_dims(self) -> int: return self.first_reduce-self.local_dims - + def global_dims(self) -> int: return sum([1 for x in self.axis_types if x == AxisType.GLOBAL]) if hasattr(self, 'axis_types') else 0 @property - def local_dims(self) -> int: return self.info.local_dims - + def local_dims(self) -> int: return sum([1 for x in self.axis_types if x == AxisType.LOCAL]) if hasattr(self, 'axis_types') else 0 @property - def upcasted(self) -> int: return self.info.upcasted - + def upcasted(self) -> int: return sum([1 for x in self.axis_types if x in {AxisType.UPCAST, AxisType.UNROLL}]) if hasattr(self, 'axis_types') else 0 @property - def dont_use_locals(self) -> bool: return self.info.dont_use_locals - - @property - def applied_opts(self) -> list[Opt]: return list(self.info.applied_opts) + def group_for_reduces(self) -> int: return sum([1 for x in self.axis_types if x == AxisType.GROUP_REDUCE]) if hasattr(self, 'axis_types') else 0 def _legacy_colors(self) -> list[str]: # first non local non reduce dims are global (blue) @@ -193,11 +185,6 @@ class Kernel: def permute(st:ShapeTracker): return st.permute(tuple(axis)) if axis is not None else st self.sts = [permute(reshape(st)) for st in self.sts] - # drops the final dimension - def upcast(self): - check(self.full_shape[-1] != 1, "can't upcast a dimension with size 1") - self.update_info(upcasted=self.info.upcasted + 1) - # axis : the axis to pull from # amount : the amount to take # top : if you want to pull that amount from the top @@ -218,8 +205,6 @@ class Kernel: if any(all_ones:=[s==1 for s in self.full_shape]): if hasattr(self, 'axis_types'): self.axis_types = [x for i,x in enumerate(self.axis_types) if not all_ones[i]] - self.update_info(local_dims=self.local_dims - sum(all_ones[self.first_reduce-self.local_dims:self.first_reduce]), - upcasted=self.upcasted - sum(all_ones[self.first_upcast:])) # TODO: no necessary since upcasted axis can't be un-upcasted self.reshape_and_permute(lambda shape: [x for i,x in enumerate(shape) if not all_ones[i]], None) return True return False @@ -375,7 +360,7 @@ class Kernel: check(0 <= (tc_opt:=cast(tuple, opt.arg)[1]) <= 2, "tensor core opts must have valid tc_opt") check(0 < (use_tensor_cores:=cast(tuple, opt.arg)[2]) <= 2, "use_tensor_cores value is not valid") check(self._apply_tc_opt(use_tensor_cores, cast(int, opt.axis), tc_select, tc_opt), "no tensor core available") - self.update_info(applied_opts=self.info.applied_opts + (opt,)) + self.applied_opts.append(opt) return axis = self.real_axis(opt) @@ -403,35 +388,25 @@ class Kernel: check(self.opts.has_local, "target does not support local") check(axis < self.global_dims, "local is for globals") self.shift_to(axis, amt, AxisType.LOCAL, insert_before=self.first_reduce) - self.update_info(local_dims=self.info.local_dims + 1) elif opt.op in {OptOps.GROUP, OptOps.GROUPTOP}: # green check(self.opts.has_local and self.opts.has_shared, "target does not support local or shared mem") check(self.first_reduce + self.group_for_reduces <= axis < self.first_upcast, "must be reduce axis to group") check(not self.tensor_core, "can't group with tensor cores") check(len(reduce_axes:=[i for r in self.reduceops for i in r.axis_arg]) == len(set(reduce_axes)), "can't group with parallel reduces") self.shift_to(axis, amt, AxisType.GROUP_REDUCE, top=(opt.op is OptOps.GROUPTOP), insert_before=self.first_reduce + self.group_for_reduces) - self.group_for_reduces += 1 elif opt.op is OptOps.UNROLL: # purple check(axis < self.first_upcast, "can't upcasted already upcasted") check(amt <= 32, "don't unroll more than 32") - # TODO: fix upcast_count to put purples before yellows. broken because of METAL tensor cores - #upcast_count = sum(x == y for x,y in zip(self.full_shape[-self.upcasted:], self.output_shape[-self.upcasted:])) if self.upcasted else 0 - #self.shift_to(axis, amt, insert_before=None if upcast_count == 0 else self.shape_len-upcast_count) - # first_reduce will ++, so offset loss in simplify_ones - if self.full_shape[axis] == amt and axis == self.first_reduce: self.update_info(local_dims=self.local_dims + 1) - if self.full_shape[axis] == amt and axis < self.first_reduce+self.group_for_reduces: self.group_for_reduces -= 1 # fully unrolling a GROUP self.shift_to(axis, amt, AxisType.UNROLL, insert_before=None) - self.upcast() elif opt.op is OptOps.UPCAST: # yellow check(axis < self.first_reduce, "upcast is for non-reduce") check(not (self.tensor_core and self.global_dims <= axis < self.global_dims+len(self.tensor_core.get_local_axes())), "can't upcast TC locals") check((self.opts is not None and self.opts.device == "DSP") or amt <= 16, "don't upcast more than 16") self.shift_to(axis, amt, AxisType.UPCAST, insert_before=None) - self.upcast() elif opt.op is OptOps.NOLOCALS: check(self.opts.has_local and not self.dont_use_locals, "NOLOCALS is meaningless if target does not support local or already not using locals") check(self.local_dims == 0 and self.group_for_reduces == 0, "can't have no locals with locals") - self.update_info(dont_use_locals=True) + self.dont_use_locals = True elif opt.op is OptOps.SWAP: check(axis < amt < self.global_dims, f"swap is only for globals with axis < amt, getting {amt=}, {axis=}, {self.global_dims=}") permute = list(range(self.shape_len)) @@ -452,7 +427,7 @@ class Kernel: padded = True check(padded, "nothing was padded") - if append_opt: self.update_info(applied_opts=self.info.applied_opts + (opt,)) + if append_opt: self.applied_opts.append(opt) if self.simplify_ones() and self.tensor_core_opts: self.tensor_core_opts.fix_axes(axis) # fix up axes in TC opts if required after simplify_ones() @@ -487,8 +462,9 @@ class Kernel: return ret.replace(src=(ret.src[0].replace(arg=st),)+ret.src[1:]) if op.op is Ops.SINK: # NOTE: should group_for_reduces be added to the local_dims? - return ret.replace(arg=replace(self.info, name=ret.arg.name if ret.arg is not None else self.name if name_override is None else name_override, - global_dims=self.global_dims if self.opts.has_local else 0, local_dims=self.local_dims + self.group_for_reduces, opts_to_apply=None)) + return ret.replace(arg=KernelInfo(ret.arg.name if ret.arg is not None else self.name if name_override is None else name_override, + self.global_dims if self.opts.has_local else 0, self.local_dims + self.group_for_reduces, + self.upcasted, self.dont_use_locals, tuple(self.applied_opts))) if op.op is Ops.REDUCE_AXIS: reduce_idx = len(self.bufs) + self.reduceops.index(op) * 2