From 0df4355cd8254ddc83701918a875c4afdd7ec593 Mon Sep 17 00:00:00 2001 From: George Hotz Date: Sun, 29 Jun 2025 09:04:37 -0700 Subject: [PATCH] upcasted warp experiments --- extra/upcasted_warps.py | 24 ++++++++++++++++++++++++ tinygrad/codegen/__init__.py | 17 +++++++++++++++++ 2 files changed, 41 insertions(+) create mode 100644 extra/upcasted_warps.py diff --git a/extra/upcasted_warps.py b/extra/upcasted_warps.py new file mode 100644 index 0000000000..f4a980a5a7 --- /dev/null +++ b/extra/upcasted_warps.py @@ -0,0 +1,24 @@ +# play with upcasted warps +from tinygrad import Tensor, Device +from tinygrad.uop.ops import KernelInfo +from tinygrad.opt import get_optimized_ast +from tinygrad.opt.kernel import OptOps, Opt +from tinygrad.engine.realize import get_program + +if __name__ == "__main__": + renderer = Device.default.renderer + N = 64 + a = Tensor.empty(N,N) + b = Tensor.empty(N,N) + # metal TC + #opts = (Opt(OptOps.UPCAST, 0, 2), # not the warp + # Opt(OptOps.UPCAST, 0, 2), Opt(OptOps.UPCAST, 1, 2), Opt(OptOps.UPCAST, 1, 2), + # Opt(OptOps.UPCAST, 0, 2), Opt(OptOps.UPCAST, 1, 2)) + # new TC should just be able to extract from this and swizzle as needed + opts = (Opt(OptOps.UPCAST, 0, 8), Opt(OptOps.UPCAST, 1, 8), Opt(OptOps.UNROLL, 0, 8)) + c = (a@b) + ast = c.schedule()[-1].ast + ast = ast.replace(arg=KernelInfo(opts_to_apply=opts)) + ast = get_optimized_ast(ast, renderer) + prg = get_program(ast, renderer) + print(prg.src) diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index ea76d43e17..7537abc54d 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -37,6 +37,17 @@ def get_rewrites_for_renderer(opts:Renderer, linearizer:bool=True) -> list[Rewri # cache with the values of the context vars return _get_rewrites_for_renderer(opts, linearizer, QUANTIZE.value, DEVECTORIZE.value, TRANSCENDENTAL.value) +# tensor cores + +from tinygrad.uop.ops import PatternMatcher, UPat, UOp + +def tensor_cores(a:UOp, b:UOp, r:UOp): + print("use tensor cores") + +pm_tensor_cores = PatternMatcher([ + ((UPat.var().gep(name='a') * UPat.var().gep(name='b')).reduce(name='r', allow_any_len=True), tensor_cores), +]) + @functools.cache def _get_rewrites_for_renderer(opts:Renderer, linearizer:bool, _QUANTIZE, _DEVECTORIZE, _TRANSCENDENTAL) -> list[RewriteStep]: # ** lowerer (rewrite_shapetracker_with_index) ** @@ -50,10 +61,16 @@ def _get_rewrites_for_renderer(opts:Renderer, linearizer:bool, _QUANTIZE, _DEVEC # expand ret.append(RewriteStep(sym+expander, name="expander")) + # use tensor cores + ret.append(RewriteStep(pm_tensor_cores, name="tensor cores")) + # ** devectorizer (full_graph_rewrite) ** # remove reduce ret.append(RewriteStep(pm_reduce+gep_pushing, lambda _: ReduceContext(), name="remove_reduce")) + # factorize warp (before gpu dims) + #ret.append(RewriteStep(pm_warp, name="warpcast")) + # add gpu dims (late) ret.append(RewriteStep(pm_add_gpudims, lambda _: opts, name="add gpudims"))