delete info from kernel [pr] (#11139)

* delete info from kernel [pr]

* update kernel info

* delete info
This commit is contained in:
George Hotz
2025-07-08 15:53:13 -07:00
committed by GitHub
parent 359bed74f8
commit a1b8f3e64f
+14 -38
View File
@@ -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