forked from tinygrad/tinygrad
add finalized to kernel [pr] (#11132)
* add finalized to kernel [pr] * add copy
This commit is contained in:
@@ -65,6 +65,9 @@ class Kernel:
|
||||
self.tensor_core_opts: Optional[TensorCoreOptions] = None
|
||||
self.use_tensor_cores: int = 0
|
||||
|
||||
# finalized means you can't optimize anymore
|
||||
self.finalized: bool = False
|
||||
|
||||
# group simplifies
|
||||
self.simplify_ones()
|
||||
self.simplify_merge_adjacent()
|
||||
@@ -87,6 +90,7 @@ class Kernel:
|
||||
# parameters for optimizations
|
||||
ret.info, ret.group_for_reduces = self.info, self.group_for_reduces
|
||||
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
|
||||
|
||||
@@ -343,6 +347,7 @@ class Kernel:
|
||||
return opt.axis
|
||||
|
||||
def apply_opt(self, opt:Opt, append_opt:bool=True):
|
||||
if self.finalized: raise RuntimeError("can't optimize Kernel after it's finalized")
|
||||
if self.dont_use_locals: check(opt.op not in {OptOps.LOCAL, OptOps.GROUP, OptOps.GROUPTOP}, "not using locals")
|
||||
|
||||
if opt.op is OptOps.TC:
|
||||
@@ -524,6 +529,7 @@ class Kernel:
|
||||
return UOp(Ops.LOAD, op.dtype, (local_buffer.view(st), UOp.store(local_buffer.view(st), grouped_reduce)))
|
||||
|
||||
return ret
|
||||
self.finalized = True
|
||||
fixed_ast = fixup_ast(self.ast)
|
||||
del fixup_ast
|
||||
return graph_rewrite(fixed_ast, view_left, name="fixup optimized AST")
|
||||
|
||||
@@ -64,7 +64,7 @@ def _try_compile_linearized_w_idx(x:tuple[int,Kernel], compiler:Compiler) -> tup
|
||||
signal.alarm(getenv("BEAM_TIMEOUT_SEC", 10))
|
||||
ret = None
|
||||
try:
|
||||
p = x[1].to_program(name_override="test")
|
||||
p = x[1].copy().to_program(name_override="test")
|
||||
assert p.uops is not None, "uop list wasn't generated?"
|
||||
if len(p.uops) >= (uops_max:=getenv("BEAM_UOPS_MAX", 3000)) > 0:
|
||||
if getenv("BEAM_LOG_SURPASS_MAX"): print(f"too many uops. {len(p.uops)=}, {uops_max=}")
|
||||
|
||||
Reference in New Issue
Block a user