From 7b865ed03d314dc73debd6ffc2975218fbe6c4a4 Mon Sep 17 00:00:00 2001 From: Andrey Date: Wed, 26 Mar 2025 08:36:48 -0300 Subject: [PATCH 1/4] use tuple in isinstance for type checking (#9583) --- examples/so_vits_svc.py | 2 +- extra/export_model.py | 2 +- extra/sqtt/rgptool.py | 2 +- extra/torch_hook/hook_torch.py | 2 +- tinygrad/ops.py | 2 +- 5 files changed, 5 insertions(+), 5 deletions(-) diff --git a/examples/so_vits_svc.py b/examples/so_vits_svc.py index 2d443cde2c..95e90fa696 100644 --- a/examples/so_vits_svc.py +++ b/examples/so_vits_svc.py @@ -504,7 +504,7 @@ def load_checkpoint_enc(checkpoint_path, model: ContentVec, optimizer=None, skip obj, v = getattr(parent, "weight"), weight_norm(weight_v, weight_g, 0) weight_g, weight_v, parent, skip = None, None, None, False if not skip and obj.shape == v.shape: - if "feature_extractor" in key and (isinstance(parent, nn.GroupNorm) or isinstance(parent, nn.LayerNorm)): # cast + if "feature_extractor" in key and (isinstance(parent, (nn.GroupNorm, nn.LayerNorm))): # cast obj.assign(v.to(obj.device).float()) else: obj.assign(v.to(obj.device)) diff --git a/extra/export_model.py b/extra/export_model.py index 8e60373f29..4f326432e9 100644 --- a/extra/export_model.py +++ b/extra/export_model.py @@ -38,7 +38,7 @@ def jit_model(model, *args) -> Tuple[TinyJit,Dict[int,str]]: @TinyJit def run(*x): out = model.forward(*x) if hasattr(model, "forward") else model(*x) - assert isinstance(out, tuple) or isinstance(out, list) or isinstance(out, Tensor), "model output must be a Tensor, tuple, or a list of Tensors for export" + assert isinstance(out, (tuple, list, Tensor)), "model output must be a Tensor, tuple, or a list of Tensors for export" out = [out] if isinstance(out, Tensor) else out return [o.realize() for o in out] diff --git a/extra/sqtt/rgptool.py b/extra/sqtt/rgptool.py index 3244cc153b..a0d499e62e 100755 --- a/extra/sqtt/rgptool.py +++ b/extra/sqtt/rgptool.py @@ -22,7 +22,7 @@ CHUNK_CLASSES = { } def pretty(val, pad=0) -> str: - if isinstance(val, ctypes.Structure) or isinstance(val, ctypes.Union): + if isinstance(val, (ctypes.Structure, ctypes.Union)): nl = '\n' # old python versions don't support \ in f-strings return f"{val.__class__.__name__}({nl}{' '*(pad+2)}{(f', {nl}'+' '*(pad+2)).join([f'{field[0]}={pretty(getattr(val, field[0]), pad=pad+2)}' for field in val._fields_])}{nl}{' '*pad})" if isinstance(val, ctypes.Array): diff --git a/extra/torch_hook/hook_torch.py b/extra/torch_hook/hook_torch.py index 24aa780ed1..f3b8113b4a 100644 --- a/extra/torch_hook/hook_torch.py +++ b/extra/torch_hook/hook_torch.py @@ -39,7 +39,7 @@ class DispatchLog(TorchDispatchMode): should_call_tiny = kwargs.get('device') is not None and kwargs['device'].type == "cuda" def can_print_arg(arg): - return args is None or isinstance(arg, str) or isinstance(arg, int) or isinstance(arg, float) or isinstance(arg, bool) + return args is None or isinstance(arg, (str, int, float, bool)) def create_tiny_mapping(arg): if WRAP_TINY: diff --git a/tinygrad/ops.py b/tinygrad/ops.py index 554428f927..5173ce1828 100644 --- a/tinygrad/ops.py +++ b/tinygrad/ops.py @@ -713,7 +713,7 @@ class UPat(MathTrait): def __init__(self, op:Optional[Union[Ops, tuple[Ops, ...], set[Ops]]]=None, dtype:Optional[Union[DType, tuple[DType, ...]]]=None, src:Optional[Union[tuple[UPat, ...], list[UPat], UPat]]=None, arg:Any=None, name:Optional[str]=None, allow_any_len:bool=False, location=None, custom_early_reject:Optional[set[Ops]]=None): - assert op is None or isinstance(op, Ops) or isinstance(op, tuple) or isinstance(op, set), "op must be Ops or tuple of Ops" + assert op is None or isinstance(op, (Ops, tuple, set)), "op must be Ops or tuple of Ops" self.op: Optional[tuple[Ops, ...]] = (op,) if isinstance(op, Ops) else (tuple(op) if isinstance(op, set) else op) self.dtype: Optional[tuple[DType, ...]] = (dtype,) if isinstance(dtype, DType) else dtype self.arg, self.name, self._in_src, self.custom_early_reject = arg, name, src, custom_early_reject From e88a640ca5436cff81d781b5e71b6a4a4cb102c7 Mon Sep 17 00:00:00 2001 From: nimlgen <138685161+nimlgen@users.noreply.github.com> Date: Wed, 26 Mar 2025 18:42:43 +0700 Subject: [PATCH 2/4] fix _access_resources for offset buffers (#9580) * fix _access_resources for offset buffers * test --- test/test_graph.py | 17 +++++++++++++++++ tinygrad/engine/jit.py | 4 +++- 2 files changed, 20 insertions(+), 1 deletion(-) diff --git a/test/test_graph.py b/test/test_graph.py index d37c4c50ba..ffc04844ed 100644 --- a/test/test_graph.py +++ b/test/test_graph.py @@ -38,6 +38,10 @@ def helper_alloc_rawbuffer(device, fill=False): rawbuf.copyin(Tensor(data).realize().lazydata.base.realized.as_buffer()) return rawbuf +def helper_create_offset_rawbuffer(base, offset=0): + x = Buffer(base.device, base.size-offset, base.dtype, base=base, offset=offset) + return x.ensure_allocated() + def helper_run_jit(jis, bufs, out_buffers): for rawbuf in out_buffers: mv = memoryview(bytearray(rawbuf.size * rawbuf.dtype.itemsize)) @@ -229,5 +233,18 @@ class TestGraph(unittest.TestCase): helper_test_graphs(Device[d0].graph, graphs) + def test_graph_offset_bufs(self): + d0 = Device.DEFAULT + if not hasattr(Device[d0].allocator, "_offset"): self.skipTest("device does not support _offset") + + b0 = [helper_alloc_rawbuffer(d0, fill=True) for _ in range(1)] + b0 += [helper_create_offset_rawbuffer(b0[0]), helper_create_offset_rawbuffer(b0[0])] + + graphs = [ + [helper_copy_op(d0, b0[0], b0[2]), helper_exec_op(d0, b0[1], [b0[0], b0[2]])], + ] + + helper_test_graphs(Device[d0].graph, graphs) + if __name__ == '__main__': unittest.main() diff --git a/tinygrad/engine/jit.py b/tinygrad/engine/jit.py index 98e9d43e2d..1400d8d772 100644 --- a/tinygrad/engine/jit.py +++ b/tinygrad/engine/jit.py @@ -120,7 +120,9 @@ class GraphRunner(Runner): if id(rawbuf.base._buf) in self.w_dependency_map: wait_nodes.append(self.w_dependency_map[id(rawbuf.base._buf)]) if i in write: if id(rawbuf.base._buf) in self.r_dependency_map: wait_nodes.extend(self.r_dependency_map.pop(id(rawbuf.base._buf))) - self.w_dependency_map[id(rawbuf.base._buf)] = new_dependency + + for i,rawbuf in enumerate(rawbufs): + if i in write: self.w_dependency_map[id(rawbuf.base._buf)] = new_dependency else: self.r_dependency_map[id(rawbuf.base._buf)].append(new_dependency) return list({id(x):x for x in wait_nodes}.values()) From 1e6e75e39a995cf831494508815a920ed55a2696 Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Wed, 26 Mar 2025 20:01:21 +0800 Subject: [PATCH 3/4] little changes from dsp branch (#9582) * little changes from dsp branch * not that one * need the where * Revert "need the where" This reverts commit 140f89c878c4b19203767cebffa209d9014c1ed1. --- extra/replay_pkl.py | 146 +++++++---------------------------- tinygrad/codegen/expander.py | 5 +- tinygrad/codegen/kernel.py | 3 +- 3 files changed, 34 insertions(+), 120 deletions(-) diff --git a/extra/replay_pkl.py b/extra/replay_pkl.py index abb8a1a33b..01ba858572 100644 --- a/extra/replay_pkl.py +++ b/extra/replay_pkl.py @@ -22,131 +22,41 @@ if __name__ == "__main__": if knum == (pknum:=getenv("KNUM", 0)) or pknum == 0: p: ProgramSpec = ei.prg.p k = Kernel(p.ast, Device["DSP"].renderer) - dsp_bufs = [Buffer("DSP", 1024+b.size*2, b.dtype).view(b.size, b.dtype, 512) for b in ei.bufs] + dsp_bufs = [Buffer("DSP", 8192+b.size, b.dtype).view(b.size, b.dtype, 4096) for b in ei.bufs] if BEAM: from tinygrad.engine.search import beam_search k = beam_search(k, dsp_bufs, BEAM.value, bool(getenv("BEAM_ESTIMATE", 1))) elif not getenv("NOOPT"): - # only NCHW - """ - if knum in [6,7,9,11]: - k.apply_opt(Opt(OptOps.PADTO, 1, 128)) - k.apply_opt(Opt(OptOps.UPCAST, 1, 128)) - elif knum in [5,8]: - k.apply_opt(Opt(op=OptOps.UNROLL, axis=1, arg=0)) - k.apply_opt(Opt(op=OptOps.UNROLL, axis=0, arg=0)) - k.apply_opt(Opt(OptOps.PADTO, 2, 128)) - k.apply_opt(Opt(OptOps.UPCAST, 2, 128)) - elif knum == 2: - k.apply_opt(Opt(op=OptOps.UNROLL, axis=1, arg=0)) - k.apply_opt(Opt(op=OptOps.UNROLL, axis=0, arg=0)) - k.apply_opt(Opt(OptOps.PADTO, 2, 128)) - k.apply_opt(Opt(OptOps.UPCAST, 2, 128)) - #k.apply_opt(Opt(op=OptOps.UPCAST, axis=1, arg=4)) - elif knum == 1: - k.apply_opt(Opt(op=OptOps.UNROLL, axis=2, arg=0)) - k.apply_opt(Opt(op=OptOps.UNROLL, axis=1, arg=0)) - #k.apply_opt(Opt(op=OptOps.UNROLL, axis=0, arg=0)) - k.apply_opt(Opt(OptOps.PADTO, 2, 128)) - k.apply_opt(Opt(OptOps.UPCAST, 2, 128)) - elif knum == 3: - k.apply_opt(Opt(op=OptOps.UNROLL, axis=0, arg=4)) - k.apply_opt(Opt(OptOps.UPCAST, 1, 128)) - elif knum == 29: - #k.apply_opt(Opt(OptOps.UPCAST, 1, 2)) - k.apply_opt(Opt(OptOps.PADTO, 1, 128)) - k.apply_opt(Opt(OptOps.UPCAST, 1, 256)) - #k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) - else: - k.hand_coded_optimizations() - """ - """ - if knum == 3: - # 12544x32 * 32x16 -> 12544x16 - + if knum == 1: + k.apply_opt(Opt(OptOps.UPCAST, 2, 32)) + k.apply_opt(Opt(OptOps.UPCAST, 1, 4)) + elif knum == 66: + k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) + k.apply_opt(Opt(OptOps.UPCAST, 0, 8)) + elif k.full_shape[-3:] == (32,3,3): + #if k.full_shape[-4]%4 != 0: k.apply_opt(Opt(OptOps.PADTO, len(k.full_shape)-4, 4)) + # 3x3 dwconv k.apply_opt(Opt(OptOps.UNROLL, 0, 0)) - k.apply_opt(Opt(OptOps.UPCAST, 1, 16)) - k.apply_opt(Opt(OptOps.UPCAST, 0, 128//16)) - #k.apply_opt(Opt(OptOps.UPCAST, 0, 256//16)) - #k.apply_opt(Opt(OptOps.UPCAST, 0, 8)) - pass - elif knum == 6: - k.apply_opt(Opt(OptOps.UNROLL, 0, 8)) - k.apply_opt(Opt(OptOps.UPCAST, 1, 0)) - elif knum == 4: - # 12544x16 * 16x96 -> 12544x96 - # (with the biased add) - #k.apply_opt(Opt(OptOps.UPCAST, 1, 96)) - #k.apply_opt(Opt(OptOps.UPCAST, 0, 4)) - #k.apply_opt(Opt(OptOps.UNROLL, 0, 0)) - #k.apply_opt(Opt(OptOps.PADTO, 0, 3)) - pass - elif knum == 13: - # 784x144 * 144x32 -> 784x32 - #k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) - #k.apply_opt(Opt(OptOps.UNROLL, 0, 2)) - #k.apply_opt(Opt(OptOps.UPCAST, 0, 4)) - #k.apply_opt(Opt(OptOps.UPCAST, 0, 2)) - #k.apply_opt(Opt(OptOps.UPCAST, 1, 32)) - pass - elif knum == 20: - # 784x192 * 192x32 -> 784x32 + k.apply_opt(Opt(OptOps.UNROLL, 0, 0)) + k.apply_opt(Opt(OptOps.UPCAST, len(k.full_shape)-3, 32)) + if k.full_shape[-4]%4 == 0: k.apply_opt(Opt(OptOps.UPCAST, len(k.full_shape)-4, 4)) + elif len(k.full_shape) == 3 and k.full_shape[1] == 32: + #if k.full_shape[0]%4 != 0: k.apply_opt(Opt(OptOps.PADTO, 0, 4)) + # weight without more k.apply_opt(Opt(OptOps.UNROLL, 0, 8)) k.apply_opt(Opt(OptOps.UPCAST, 1, 32)) - k.apply_opt(Opt(OptOps.UPCAST, 0, 4)) - elif knum == 35: - k.apply_opt(Opt(OptOps.UNROLL, 0, 128)) - k.apply_opt(Opt(OptOps.UPCAST, 0, 2)) - k.apply_opt(Opt(OptOps.UPCAST, 1, 64)) - elif knum == 37: - pass - elif knum == 24: - #k.apply_opt(Opt(OptOps.UNROLL, 0, 0)) - k.apply_opt(Opt(OptOps.UPCAST, 1, 64)) - #k.apply_opt(Opt(OptOps.UPCAST, 0, 2)) - """ - #if knum in [7, 11, 14, 18]: - # alignment issue? - #pass - if knum == 4: - k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) - k.apply_opt(Opt(OptOps.UPCAST, 1, 96)) - k.apply_opt(Opt(OptOps.UPCAST, 0, 4)) - elif knum == 6: - k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) - k.apply_opt(Opt(OptOps.UPCAST, 1, 24)) - k.apply_opt(Opt(OptOps.UPCAST, 0, 16)) - elif knum == 11: - k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) - k.apply_opt(Opt(OptOps.UPCAST, 1, 144)) - #k.apply_opt(Opt(OptOps.UPCAST, 0, 8)) - elif knum == 14: - k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) - k.apply_opt(Opt(OptOps.UPCAST, 1, 192)) - k.apply_opt(Opt(OptOps.UPCAST, 0, 2)) - elif knum == 37: - k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) - k.apply_opt(Opt(OptOps.UPCAST, 1, 384)) - else: - full_shape = k.full_shape - out_shape = k.sts[0].shape - out_strides = k.sts[0].real_strides() - if len(out_strides) == 3: - if full_shape[1] < 128: - if full_shape[2] <= 32: k.apply_opt(Opt(OptOps.UNROLL, 0, 0)) - else: k.apply_opt(Opt(OptOps.UNROLL, 0, 8)) - k.apply_opt(Opt(OptOps.UPCAST, 1, full_shape[1])) - if out_strides[0] < 128: - upcast_0 = 128//out_strides[0] - if out_shape[0]%upcast_0 == 0 and upcast_0 != 1: k.apply_opt(Opt(OptOps.UPCAST, 0, upcast_0)) - elif full_shape[1] % 128 == 0: - k.apply_opt(Opt(OptOps.UPCAST, 1, 128)) - elif len(out_strides) == 1: - #if full_shape[0]%128 == 0: k.apply_opt(Opt(OptOps.UPCAST, 0, 128)) - pass - #print("here", out_shape, out_strides, k.name) - #k.hand_coded_optimizations() - #if knum in [5]: k.apply_opt(Opt(OptOps.UPCAST, 1, 2)) + if k.full_shape[0]%4 == 0: k.apply_opt(Opt(OptOps.UPCAST, 0, 4)) + elif len(k.full_shape) == 4 and k.full_shape[2] == 32: + #if k.full_shape[1]%4 != 0: k.apply_opt(Opt(OptOps.PADTO, 1, 4)) + # weight with more + k.apply_opt(Opt(OptOps.UNROLL, 0, 8)) + k.apply_opt(Opt(OptOps.UPCAST, 2, 32)) + if k.full_shape[1]%4 == 0: k.apply_opt(Opt(OptOps.UPCAST, 1, 4)) + elif len(k.full_shape) == 1: + for sz in [128,64,32]: + if k.full_shape[0]%sz == 0: + k.apply_opt(Opt(OptOps.UPCAST, 0, sz)) + break p2 = k.to_program() new_ei = replace(ei, prg=CompiledRunner(p2), bufs=dsp_bufs) new_ei.run() diff --git a/tinygrad/codegen/expander.py b/tinygrad/codegen/expander.py index 0175d6475c..8585a7e92e 100644 --- a/tinygrad/codegen/expander.py +++ b/tinygrad/codegen/expander.py @@ -50,6 +50,9 @@ def do_expand(root:UOp): if root.op is Ops.IF: # for the first arg of IF, just pass them through ignoring UNROLLS new_srcs.append(src) + elif root.op is Ops.REDUCE and src.op is Ops.RANGE: + # for any range args of REDUCE, pass them through + new_srcs.append(src) elif src.dtype.count > 1: # put any input dtype > 1 grouped together new_srcs.append(UOp(Ops.CAT, src.dtype.scalar().vec(expand_sz*src.dtype.count), (src,)*expand_sz)) @@ -82,7 +85,7 @@ expander = PatternMatcher([ lambda outer, inner: UOp(Ops.UNROLL, outer.dtype, (inner.src[0],), inner.arg+outer.arg)), # do expansion (UPat((*GroupOp.ALU, Ops.CAST, Ops.BITCAST, Ops.GEP, Ops.WMMA, Ops.LOAD, Ops.STORE, Ops.INDEX, Ops.ASSIGN, - Ops.VECTORIZE, Ops.IF), name="root", custom_early_reject=set([Ops.UNROLL])), do_expand), + Ops.VECTORIZE, Ops.IF, Ops.REDUCE), name="root", custom_early_reject=set([Ops.UNROLL])), do_expand), (UPat(Ops.CONTRACT, name="con"), do_contract), # vectorize DEFINE_ACC (UPat(Ops.VECTORIZE, src=UPat(Ops.DEFINE_ACC, name="acc"), name="v"), diff --git a/tinygrad/codegen/kernel.py b/tinygrad/codegen/kernel.py index 9143458f25..e0085e498e 100644 --- a/tinygrad/codegen/kernel.py +++ b/tinygrad/codegen/kernel.py @@ -669,7 +669,8 @@ class Kernel: print(self.name) if DEBUG >= 5: print(self.ast) for i,(buf,st) in enumerate([(buf,st) for buf,st in zip(self.bufs, self.sts) if buf.op not in {Ops.CONST, Ops.VALID}]): - print(f"{i:2d}: {str(st.shape):25s} {str(buf.src[0].dtype).replace('dtypes.',''):20s}", st.real_strides()) + print(f"{i:2d}: {str(st.shape):25s} {str(buf.src[0].dtype).replace('dtypes.',''):20s} {str(st.real_strides()):30s}", + str(st) if DEBUG >= 4 else "") print(self.applied_opts) if DEBUG >= 5: print(modified_ast) # verify AST matches the spec after applying opts From 5c6cd884e3fd0782b6daedcaff48c5d2ae04708b Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Wed, 26 Mar 2025 21:42:52 +0800 Subject: [PATCH 4/4] multiple simplifies is faster [pr] (#9586) * multiple simplifies is faster [pr] * cleanup * cleanup --- tinygrad/codegen/devectorizer.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/tinygrad/codegen/devectorizer.py b/tinygrad/codegen/devectorizer.py index d2096bde83..c7b95fd9fd 100644 --- a/tinygrad/codegen/devectorizer.py +++ b/tinygrad/codegen/devectorizer.py @@ -4,7 +4,7 @@ from collections import defaultdict from tinygrad.dtype import dtypes, ImageDType, PtrDType from tinygrad.ops import UOp, Ops, UPat, PatternMatcher, resolve from tinygrad.ops import graph_rewrite, GroupOp -from tinygrad.codegen.symbolic import symbolic_simple, split_uop, uop_given_valid, parse_valid, simplify_valid, sym +from tinygrad.codegen.symbolic import symbolic_simple, split_uop, uop_given_valid, parse_valid, simplify_valid, sym, symbolic from tinygrad.helpers import getenv, flatten, TRANSCENDENTAL, AMX, prod, DEVECTORIZE from tinygrad.codegen.transcendental import xexp2, xlog2, xsin, xpow, TRANSCENDENTAL_SUPPORTED_DTYPES from tinygrad.renderer import Renderer @@ -15,13 +15,16 @@ def expand_index(buf:UOp, vec:UOp, mask:UOp|None=None): if getenv("UNSAFE_DISABLE_MASK", 0): mask = None # first, extract all the relevant offsets offsets_rootsrc: defaultdict[Any, dict[int, list[int]]] = defaultdict(dict) + midx, mmask = graph_rewrite(UOp.sink(UOp.sink(*[vec.gep(i) for i in range(vec.dtype.count)]), + UOp.sink(*[mask.gep(i) for i in range(vec.dtype.count)]) if mask is not None else UOp(Ops.NOOP)), + symbolic, name=f"index_buf_{buf.arg}").src for i in range(vec.dtype.count): - idx = vec.gep(i).simplify() + idx: Any = midx.src[i] if idx.op is Ops.ADD and idx.src[1].op is Ops.CONST: root_src, arg = idx.src[0], idx.src[1].arg elif idx.op is Ops.ADD and idx.src[0].op is Ops.CONST: root_src, arg = idx.src[1], idx.src[0].arg elif idx.op is Ops.CONST: root_src, arg = "CONST", idx.arg else: root_src, arg = idx, 0 - if mask is not None: root_src = (mask.gep(i).simplify(), root_src) + if mask is not None: root_src = (mmask.src[i], root_src) offsets_rootsrc[root_src].setdefault(arg, []).append(i) # the buf.dtype is always a pointer