diff --git a/tinygrad/codegen/opt/search.py b/tinygrad/codegen/opt/search.py index c2dc093e67..20e45c9a46 100644 --- a/tinygrad/codegen/opt/search.py +++ b/tinygrad/codegen/opt/search.py @@ -1,9 +1,8 @@ -from typing import cast import functools, math, time, multiprocessing, traceback, signal, atexit from dataclasses import replace from tinygrad.uop.ops import sym_infer, AxisType, pyrender from tinygrad.device import Device, Buffer, Compiler -from tinygrad.helpers import prod, flatten, DEBUG, CACHELEVEL, diskcache_get, diskcache_put, getenv, Context, colored, time_to_str +from tinygrad.helpers import prod, flatten, DEBUG, CACHELEVEL, diskcache_get, diskcache_put, getenv, Context, colored, time_to_str, unwrap from tinygrad.helpers import IGNORE_BEAM_CACHE from tinygrad.codegen.opt import Opt, OptOps, KernelOptError from tinygrad.tensor import Tensor @@ -50,7 +49,7 @@ def _time_program(p:ProgramSpec, lib:bytes, var_vals:dict[str, int], rawbufs:lis if hasattr(dev:=Device[p.device], 'invalidate_caches'): dev.invalidate_caches() else: with Context(DEBUG=0, BEAM=0, CAPTURING=0, TRACK_MATCH_STATS=0): Tensor.ones(1024,1024).contiguous().realize(do_update_stats=False) - tms.append(cast(float, car(input_bufs, var_vals, wait=True))*factor) + tms.append(unwrap(car(input_bufs, var_vals, wait=True))*factor) if early_stop is not None and early_stop < min(tms): break return tms @@ -168,7 +167,7 @@ def beam_search(s:Scheduler, rawbufs:list[Buffer], amt:int, allow_test_size=True raise timed.append((candidates[i], min(tms))) if BEAM_DEBUG > 1: - print(f"{time.perf_counter() - st:7.2f}s: {i:5d} {len(cast(list, p.uops)):5d} uops", + print(f"{time.perf_counter() - st:7.2f}s: {i:5d} {len(unwrap(p.uops)):5d} uops", f"{time_to_str(compile_et, w=12)} compile/{time_to_str(timed[-1][1], w=12)} run", f" {len(timed):4d}/{len(candidates):4d} {timed[-1][0].colored_shape()}") elif DEBUG >= 2: diff --git a/tinygrad/engine/realize.py b/tinygrad/engine/realize.py index 50458a533a..177f327380 100644 --- a/tinygrad/engine/realize.py +++ b/tinygrad/engine/realize.py @@ -3,6 +3,7 @@ import time, pprint, random, itertools, math from dataclasses import dataclass, replace, field from tinygrad.helpers import all_same, colored, DEBUG, GlobalCounters, ansilen, BEAM, NOOPT, all_int, CAPTURING, Metadata, TRACEMETA, TracingKey from tinygrad.helpers import DEVECTORIZE, time_to_str, VALIDATE_WITH_CPU, getenv, cpu_profile, PROFILE, ProfilePointEvent, cpu_events, prod, Context +from tinygrad.helpers import unwrap from tinygrad.uop.ops import Ops, PatternMatcher, UOp, UPat, sym_infer, graph_rewrite, print_uops, track_rewrites, KernelInfo, pyrender from tinygrad.device import Device, Buffer from tinygrad.renderer import Renderer, ProgramSpec, Estimates @@ -165,7 +166,7 @@ class ExecItem: fixedvars: dict[str, int] = field(default_factory=dict) def run(self, _var_vals:dict[str, int]|None=None, wait=False, jit=False, do_update_stats=True) -> float|None: var_vals = self.fixedvars if _var_vals is None else (_var_vals|self.fixedvars) - bufs = [cast(Buffer, x) for x in self.bufs] if jit else [cast(Buffer, x).ensure_allocated() for x in self.bufs] + bufs = [unwrap(x) for x in self.bufs] if jit else [unwrap(x).ensure_allocated() for x in self.bufs] if PROFILE: payload = {"metadata":self.metadata, "var_vals":var_vals, "bufs":[b.trace_num for b in bufs], "name":self.prg.display_name} payload["outputs"], payload["inputs"] = (self.prg.p.outs, self.prg.p.ins) if isinstance(self.prg, CompiledRunner) else ([0], [1]) diff --git a/tinygrad/runtime/ops_null.py b/tinygrad/runtime/ops_null.py index 07f5494ca7..5ff75b2da9 100644 --- a/tinygrad/runtime/ops_null.py +++ b/tinygrad/runtime/ops_null.py @@ -1,5 +1,4 @@ import functools -from typing import cast from tinygrad.device import Compiled, Compiler, Allocator from tinygrad.engine.jit import MultiGraphRunner from tinygrad.renderer.cstyle import Renderer, CStyleLanguage @@ -33,7 +32,7 @@ class NullGraph(MultiGraphRunner): class NullDevice(Compiled): def __init__(self, device:str): renderer:functools.partial|type[Renderer] - match cast(str, EMULATE.value): + match str(EMULATE.value): case "AMD": renderer = functools.partial(AMDLLVMRenderer, "gfx1100") case "AMD_RDNA4": renderer = functools.partial(AMDLLVMRenderer, "gfx1201") case "": renderer = NullRenderer diff --git a/tinygrad/schedule/rangeify.py b/tinygrad/schedule/rangeify.py index fb7443a858..376689294d 100644 --- a/tinygrad/schedule/rangeify.py +++ b/tinygrad/schedule/rangeify.py @@ -1,4 +1,3 @@ -from typing import cast from dataclasses import dataclass, field import itertools from tinygrad.dtype import dtypes, PtrDType, ImageDType, AddrSpace @@ -573,5 +572,5 @@ def get_rangeify_map(sink:UOp) -> dict[UOp, UOp]: assert s.tag is not None for a in s.tag: if a is None: continue - becomes_map[uop_list[cast(int, a)]] = s.replace(tag=None) + becomes_map[uop_list[int(a)]] = s.replace(tag=None) return becomes_map