forked from tinygrad/tinygrad
restrict allowed call bodies (#17925)
Autogen / In-tree Autogen (macos) (push) Waiting to run
Benchmarks / Mac pytest (push) Waiting to run
Benchmarks / LLM (DEV=AMD) (push) Waiting to run
Benchmarks / LLM (DEV=METAL) (push) Waiting to run
Benchmarks / LLM (DEV=NV) (push) Waiting to run
Benchmarks / HLB-CIFAR10 (DEV=AMD) (push) Waiting to run
Benchmarks / HLB-CIFAR10 (DEV=METAL) (push) Waiting to run
Benchmarks / HLB-CIFAR10 (DEV=NV) (push) Waiting to run
Benchmarks / MLPerf (AMD) (push) Waiting to run
Benchmarks / MLPerf (NV) (push) Waiting to run
Benchmarks / Stable Diffusion (DEV=AMD) (push) Waiting to run
Benchmarks / Stable Diffusion (DEV=METAL) (push) Waiting to run
Benchmarks / Stable Diffusion (DEV=NV) (push) Waiting to run
Benchmarks / Multi-GPU Benchmarks (DEV=AMD) (push) Waiting to run
Benchmarks / Multi-GPU Benchmarks (DEV=NV) (push) Waiting to run
Benchmarks / Tests (DEV=AMD) (push) Waiting to run
Benchmarks / Tests (DEV=METAL) (push) Waiting to run
Benchmarks / Tests (DEV=NV) (push) Waiting to run
Benchmarks / UsbGPU Benchmark (push) Waiting to run
Benchmarks / openpilot 0.11.2 compile3 dmonitoring (DEV=QCOM) (push) Waiting to run
Benchmarks / openpilot 0.11.0 compile3 dmonitoring (DEV=QCOM) (push) Waiting to run
Benchmarks / openpilot 0.11.0 compile3 policy (DEV=QCOM) (push) Waiting to run
Benchmarks / openpilot 0.11.2 compile3 supercombo (DEV=QCOM) (push) Waiting to run
Benchmarks / openpilot 0.11.0 compile3 vision (DEV=QCOM) (push) Waiting to run
Benchmarks / openpilot 0.11.0 compile3 dmonitoring (DEV=QCOM:IR3) (push) Waiting to run
Benchmarks / openpilot 0.11.2 compile3 dmonitoring (DEV=QCOM:IR3) (push) Waiting to run
Benchmarks / openpilot 0.11.0 compile3 policy (DEV=QCOM:IR3) (push) Waiting to run
Benchmarks / openpilot 0.11.2 compile3 supercombo (DEV=QCOM:IR3) (push) Waiting to run
Benchmarks / openpilot 0.11.0 compile3 vision (DEV=QCOM:IR3) (push) Waiting to run
Benchmarks / DSP Benchmark (push) Waiting to run
Benchmarks / UsbGPU Benchmark (comma) (push) Waiting to run
Benchmarks / PCI Driver Benchmark (DEV=AMD) (push) Waiting to run
Benchmarks / PCI Driver Benchmark (DEV=NV) (push) Waiting to run
Benchmarks / LLVM Speed (push) Waiting to run
Platform Tests / MacOS (unit) (push) Waiting to run
Platform Tests / MacOS (unit, mock) (push) Waiting to run
Platform Tests / MacOS (DEV=METAL) (1) (push) Waiting to run
Platform Tests / MacOS (DEV=METAL) (2) (push) Waiting to run
Platform Tests / MacOS (DEV=CPU:CLANG) (push) Waiting to run
Platform Tests / MacOS (DEV=CPU:LLVM) (push) Waiting to run
Platform Tests / MacOS (DEV=CPU:LVP) (push) Waiting to run
Platform Tests / MacOS (DEV=WEBGPU) (push) Waiting to run
Platform Tests / Windows (DEV=CPU:CLANG) (push) Waiting to run
Platform Tests / Windows (DEV=CPU:LLVM) (push) Waiting to run
Platform Tests / Windows (DEV=CPU:X86) (push) Waiting to run
Platform Tests / Windows (DEV=WEBGPU) (push) Waiting to run
Deploy Docs / deploy (push) Successful in 4m46s
Autogen / In-tree Autogen (push) Successful in 6m58s
Unit Tests / Docs (push) Failing after 2s
Unit Tests / Torch Backend Tests (push) Failing after 0s
Unit Tests / Torch Backend Training (push) Failing after 1s
Unit Tests / Python Backend (push) Failing after 0s
Unit Tests / Linters (push) Failing after 1s
Unit Tests / Null Tests (push) Canceled after 0s
Unit Tests / Unit Tests (push) Canceled after 0s
Unit Tests / SPEC=2 (1) (push) Canceled after 0s
Unit Tests / SPEC=2 (2) (push) Canceled after 0s
Unit Tests / Fuzzing (push) Canceled after 0s
Unit Tests / CL IMAGE Tests (push) Canceled after 0s
Unit Tests / openpilot Compile Tests (push) Canceled after 0s
Unit Tests / ONNX (CPU) Tests (push) Canceled after 0s
Unit Tests / Optimization Tests (push) Canceled after 0s
Unit Tests / Test LLM (push) Canceled after 0s
Unit Tests / Models (push) Canceled after 0s
Unit Tests / Linux (DSP) (push) Canceled after 0s
Unit Tests / Linux (DEV=CL) (push) Canceled after 0s
Unit Tests / Linux (DEV=CPU:CLANG) (push) Canceled after 0s
Unit Tests / Linux (DEV=CPU:LLVM) (push) Canceled after 0s
Unit Tests / Linux (DEV=CPU:LVP) (push) Canceled after 0s
Unit Tests / Linux (DEV=CPU:X86) (push) Canceled after 0s
Unit Tests / Linux (DEV=WEBGPU) (push) Canceled after 0s
Unit Tests / AMD ASM IDE (push) Canceled after 0s
Unit Tests / Linux (am) (push) Canceled after 0s
Unit Tests / Linux (amd gfx950) (push) Canceled after 0s
Unit Tests / Linux (amd gfx1100) (push) Canceled after 0s
Unit Tests / Linux (amd gfx1201) (push) Canceled after 0s
Unit Tests / Linux (amdllvm gfx950) (push) Canceled after 0s
Unit Tests / Linux (amdllvm gfx1100) (push) Canceled after 0s
Unit Tests / Linux (amdllvm gfx1201) (push) Canceled after 0s
Unit Tests / Linux (nv) (push) Canceled after 0s
Unit Tests / Linux (ptx) (push) Canceled after 0s
Unit Tests / Compile-only (DEV=NULL:IR3:a630) (push) Canceled after 0s
Unit Tests / Compile-only (DEV=NULL:NAK:sm_120) (push) Canceled after 0s
Unit Tests / Compile-only (DEV=NULL:QCOMCL:a630) (push) Canceled after 0s
Autogen / In-tree Autogen (macos) (push) Waiting to run
Benchmarks / Mac pytest (push) Waiting to run
Benchmarks / LLM (DEV=AMD) (push) Waiting to run
Benchmarks / LLM (DEV=METAL) (push) Waiting to run
Benchmarks / LLM (DEV=NV) (push) Waiting to run
Benchmarks / HLB-CIFAR10 (DEV=AMD) (push) Waiting to run
Benchmarks / HLB-CIFAR10 (DEV=METAL) (push) Waiting to run
Benchmarks / HLB-CIFAR10 (DEV=NV) (push) Waiting to run
Benchmarks / MLPerf (AMD) (push) Waiting to run
Benchmarks / MLPerf (NV) (push) Waiting to run
Benchmarks / Stable Diffusion (DEV=AMD) (push) Waiting to run
Benchmarks / Stable Diffusion (DEV=METAL) (push) Waiting to run
Benchmarks / Stable Diffusion (DEV=NV) (push) Waiting to run
Benchmarks / Multi-GPU Benchmarks (DEV=AMD) (push) Waiting to run
Benchmarks / Multi-GPU Benchmarks (DEV=NV) (push) Waiting to run
Benchmarks / Tests (DEV=AMD) (push) Waiting to run
Benchmarks / Tests (DEV=METAL) (push) Waiting to run
Benchmarks / Tests (DEV=NV) (push) Waiting to run
Benchmarks / UsbGPU Benchmark (push) Waiting to run
Benchmarks / openpilot 0.11.2 compile3 dmonitoring (DEV=QCOM) (push) Waiting to run
Benchmarks / openpilot 0.11.0 compile3 dmonitoring (DEV=QCOM) (push) Waiting to run
Benchmarks / openpilot 0.11.0 compile3 policy (DEV=QCOM) (push) Waiting to run
Benchmarks / openpilot 0.11.2 compile3 supercombo (DEV=QCOM) (push) Waiting to run
Benchmarks / openpilot 0.11.0 compile3 vision (DEV=QCOM) (push) Waiting to run
Benchmarks / openpilot 0.11.0 compile3 dmonitoring (DEV=QCOM:IR3) (push) Waiting to run
Benchmarks / openpilot 0.11.2 compile3 dmonitoring (DEV=QCOM:IR3) (push) Waiting to run
Benchmarks / openpilot 0.11.0 compile3 policy (DEV=QCOM:IR3) (push) Waiting to run
Benchmarks / openpilot 0.11.2 compile3 supercombo (DEV=QCOM:IR3) (push) Waiting to run
Benchmarks / openpilot 0.11.0 compile3 vision (DEV=QCOM:IR3) (push) Waiting to run
Benchmarks / DSP Benchmark (push) Waiting to run
Benchmarks / UsbGPU Benchmark (comma) (push) Waiting to run
Benchmarks / PCI Driver Benchmark (DEV=AMD) (push) Waiting to run
Benchmarks / PCI Driver Benchmark (DEV=NV) (push) Waiting to run
Benchmarks / LLVM Speed (push) Waiting to run
Platform Tests / MacOS (unit) (push) Waiting to run
Platform Tests / MacOS (unit, mock) (push) Waiting to run
Platform Tests / MacOS (DEV=METAL) (1) (push) Waiting to run
Platform Tests / MacOS (DEV=METAL) (2) (push) Waiting to run
Platform Tests / MacOS (DEV=CPU:CLANG) (push) Waiting to run
Platform Tests / MacOS (DEV=CPU:LLVM) (push) Waiting to run
Platform Tests / MacOS (DEV=CPU:LVP) (push) Waiting to run
Platform Tests / MacOS (DEV=WEBGPU) (push) Waiting to run
Platform Tests / Windows (DEV=CPU:CLANG) (push) Waiting to run
Platform Tests / Windows (DEV=CPU:LLVM) (push) Waiting to run
Platform Tests / Windows (DEV=CPU:X86) (push) Waiting to run
Platform Tests / Windows (DEV=WEBGPU) (push) Waiting to run
Deploy Docs / deploy (push) Successful in 4m46s
Autogen / In-tree Autogen (push) Successful in 6m58s
Unit Tests / Docs (push) Failing after 2s
Unit Tests / Torch Backend Tests (push) Failing after 0s
Unit Tests / Torch Backend Training (push) Failing after 1s
Unit Tests / Python Backend (push) Failing after 0s
Unit Tests / Linters (push) Failing after 1s
Unit Tests / Null Tests (push) Canceled after 0s
Unit Tests / Unit Tests (push) Canceled after 0s
Unit Tests / SPEC=2 (1) (push) Canceled after 0s
Unit Tests / SPEC=2 (2) (push) Canceled after 0s
Unit Tests / Fuzzing (push) Canceled after 0s
Unit Tests / CL IMAGE Tests (push) Canceled after 0s
Unit Tests / openpilot Compile Tests (push) Canceled after 0s
Unit Tests / ONNX (CPU) Tests (push) Canceled after 0s
Unit Tests / Optimization Tests (push) Canceled after 0s
Unit Tests / Test LLM (push) Canceled after 0s
Unit Tests / Models (push) Canceled after 0s
Unit Tests / Linux (DSP) (push) Canceled after 0s
Unit Tests / Linux (DEV=CL) (push) Canceled after 0s
Unit Tests / Linux (DEV=CPU:CLANG) (push) Canceled after 0s
Unit Tests / Linux (DEV=CPU:LLVM) (push) Canceled after 0s
Unit Tests / Linux (DEV=CPU:LVP) (push) Canceled after 0s
Unit Tests / Linux (DEV=CPU:X86) (push) Canceled after 0s
Unit Tests / Linux (DEV=WEBGPU) (push) Canceled after 0s
Unit Tests / AMD ASM IDE (push) Canceled after 0s
Unit Tests / Linux (am) (push) Canceled after 0s
Unit Tests / Linux (amd gfx950) (push) Canceled after 0s
Unit Tests / Linux (amd gfx1100) (push) Canceled after 0s
Unit Tests / Linux (amd gfx1201) (push) Canceled after 0s
Unit Tests / Linux (amdllvm gfx950) (push) Canceled after 0s
Unit Tests / Linux (amdllvm gfx1100) (push) Canceled after 0s
Unit Tests / Linux (amdllvm gfx1201) (push) Canceled after 0s
Unit Tests / Linux (nv) (push) Canceled after 0s
Unit Tests / Linux (ptx) (push) Canceled after 0s
Unit Tests / Compile-only (DEV=NULL:IR3:a630) (push) Canceled after 0s
Unit Tests / Compile-only (DEV=NULL:NAK:sm_120) (push) Canceled after 0s
Unit Tests / Compile-only (DEV=NULL:QCOMCL:a630) (push) Canceled after 0s
* restrict allowed call bodies * simpler * less * simpler spec * cleanup for spec
This commit is contained in:
@@ -5,12 +5,13 @@ from tinygrad.dtype import dtypes
|
||||
from tinygrad.renderer.cstyle import CStyleLanguage
|
||||
from tinygrad.uop.ops import KernelInfo
|
||||
|
||||
# an external call is a CALL on a CUSTOM_FUNCTION body holding the callee (the loaded function pointer)
|
||||
def call_out_kernel(F:UOp, C:UOp) -> UOp:
|
||||
call = F[0].load().call(UOp.const(3).cast(dtypes.int), C[0], ret_dtype=dtypes.void)
|
||||
call = UOp.custom_function("callback", F[0].load()).call(UOp.const(3).cast(dtypes.int), C[0], ret_dtype=dtypes.void)
|
||||
return C.after(call)[1].store(C.after(call)[0].load() + 1).sink(arg=KernelInfo(name="call_out"))
|
||||
|
||||
def call_ret_kernel(F:UOp, C:UOp) -> UOp:
|
||||
val = F[0].load().call(UOp.const(21).cast(dtypes.int), ret_dtype=dtypes.int)
|
||||
val = UOp.custom_function("callback", F[0].load()).call(UOp.const(21).cast(dtypes.int), ret_dtype=dtypes.int)
|
||||
return C[0].store(val * 2).sink(arg=KernelInfo(name="call_ret"))
|
||||
|
||||
@unittest.skipUnless(isinstance(Device["CPU"].renderer, CStyleLanguage), "TODO: CALL is rendered in C style only")
|
||||
|
||||
@@ -24,7 +24,7 @@ def _make_linear(buffer_lists, copies=None):
|
||||
src0 = bufs[0].copy_to_device(bufs[1].device)
|
||||
else:
|
||||
src0 = UOp(Ops.SINK, src=tuple(bufs))
|
||||
calls.append(UOp(Ops.CALL, src=(src0, *bufs)))
|
||||
calls.append(src0.call(*bufs))
|
||||
return UOp(Ops.LINEAR, src=tuple(calls))
|
||||
|
||||
def _get_planned_view(buf:UOp) -> tuple[UOp, int, int]|None:
|
||||
|
||||
@@ -228,7 +228,7 @@ class TestViz(unittest.TestCase):
|
||||
pm = PatternMatcher([(UPat(Ops.CONST, arg=3, name="x"), lambda x: UOp.const(4, x.dtype))])
|
||||
with save_viz() as viz:
|
||||
inner = UOp.const(3)
|
||||
call = UOp(Ops.CALL, src=(UOp(Ops.SINK, src=(inner,)),))
|
||||
call = UOp.sink(inner).call()
|
||||
graph_rewrite(call, TrackedPatternMatcher(pm.patterns), enter_calls=True)
|
||||
details = list(viz.get_details(0, 0))
|
||||
self.assertTrue(details[-1]["change"], "viz replay should detect change inside CALL")
|
||||
|
||||
@@ -360,7 +360,7 @@ def full_rewrite_to_sink(ast:UOp, ren:Renderer, optimize:bool=True) -> UOp:
|
||||
# the boundary: required compute dtypes settle here; derivable const edges may stay bare
|
||||
# NOTE: we need indexing_simplify to remove the cast to long using the Invalid
|
||||
# NOTE: symbolic must NOT be composed here -- pm_data_invalid pushes the weak result CAST into a gated WHERE, remaking the weak node, and it cycles
|
||||
sink = graph_rewrite(sink, pm_lower_weak+indexing_simplify, name="lower all index dtypes")
|
||||
sink = graph_rewrite(sink, pm_lower_weak+indexing_simplify, name="lower all index dtypes", enter_calls=True)
|
||||
|
||||
# final symbolic before decomp
|
||||
sink = graph_rewrite(sink, symbolic, name="final symbolic")
|
||||
|
||||
@@ -65,9 +65,9 @@ base_rewrite = PatternMatcher([
|
||||
(UPat(GroupOp.ALU, name="x"), lambda ctx,x: ctx.code_for_op[x.op](
|
||||
*([strip_parens(ctx[v]) if v.op == x.op and x.op in {Ops.ADD, Ops.MUL, Ops.XOR, Ops.OR, Ops.AND} else ctx[v] for v in x.src]), x.dtype)),
|
||||
|
||||
# call an external function
|
||||
(UPat(Ops.CALL, src=(UPat(),), allow_any_len=True, name="x"), lambda ctx,x:
|
||||
f"((({ctx.abi}{ctx.render_dtype(x.dtype)}(*)({', '.join(ctx.render_type(y) for y in x.src[1:])}))({ctx[x.src[0]]}))" +
|
||||
# call an external function: the CUSTOM_FUNCTION body holds the callee (a function pointer), the other srcs are the args
|
||||
(UPat(Ops.CALL, src=(UPat(Ops.CUSTOM_FUNCTION, src=(UPat(name="fptr"),)),), allow_any_len=True, name="x"), lambda ctx,x,fptr:
|
||||
f"((({ctx.abi}{ctx.render_dtype(x.dtype)}(*)({', '.join(ctx.render_type(y) for y in x.src[1:])}))({ctx[fptr]}))" +
|
||||
f"({', '.join(f'({ctx.render_type(y)})({ctx[y]})' for y in x.src[1:])}))" + (";" if x.dtype is dtypes.void else "")),
|
||||
|
||||
# custom passes through with format
|
||||
@@ -214,7 +214,7 @@ class CStyleLanguage(Renderer):
|
||||
c: defaultdict[str, int] = defaultdict(int)
|
||||
name = "test"
|
||||
for u in uops:
|
||||
if u.op in {Ops.NOOP, Ops.GROUP, Ops.CONST}: continue
|
||||
if u.op in {Ops.NOOP, Ops.GROUP, Ops.CONST, Ops.CUSTOM_FUNCTION}: continue
|
||||
if u.op == Ops.STACK and len(u.src) == 0: continue
|
||||
if u.op is Ops.AFTER:
|
||||
r[u] = r[u.src[0]]
|
||||
|
||||
@@ -254,7 +254,9 @@ class USBMMIOInterface(MMIOInterface):
|
||||
def make_buf(devs, slot:int=0, tag:str="signal") -> UOp: return UOp.placeholder((1,), dtypes.uint64, slot, device=devs, volatile=True, tag=tag)
|
||||
|
||||
def _libusb(devs, dep:tuple[UOp, ...], fn:str, *args) -> UOp:
|
||||
return make_buf(devs, tag=f"func:{fn}").after(*dep).index(0).load().call(make_buf(devs, tag="usb_handle").index(0).load(),
|
||||
# the CUSTOM_FUNCTION body holds the callee (the loaded function pointer), the call args are plain dataflow
|
||||
fptr = make_buf(devs, tag=f"func:{fn}").after(*dep).index(0).load()
|
||||
return UOp.custom_function(fn, fptr).call(make_buf(devs, tag="usb_handle").index(0).load(),
|
||||
*[UOp.const(a, dtypes.int) if isinstance(a, int) else a for a in args], ret_dtype=dtypes.void)
|
||||
|
||||
def usb_bulk(devs, dep, endpoint:int, data:UOp, length, timeout:int=1000) -> UOp: # NULL actual_length out param
|
||||
|
||||
+17
-15
@@ -126,8 +126,8 @@ def dtype_from_uop(op:Ops, src:tuple[UOp,...], arg:Any) -> DType:
|
||||
# always void
|
||||
return dtypes.void
|
||||
case Ops.CALL:
|
||||
# a CALL of an opaque body (CallInfo arg) is void, a CALL of an address states its return dtype in the arg
|
||||
return arg if isinstance(arg, DType) else dtypes.void
|
||||
# a call states its (possibly void) dtype in the CallInfo
|
||||
return arg.dtype if isinstance(arg, CallInfo) else dtypes.void
|
||||
case Ops.CUSTOM | Ops.CUSTOMI:
|
||||
assert isinstance(arg, tuple) and len(arg) == 2 and isinstance(arg[1], DType), f"CUSTOM/CUSTOMI arg must be (str, DType), got {arg}"
|
||||
return arg[1]
|
||||
@@ -1202,16 +1202,19 @@ class UOp(RandMixin, metaclass=UOpMetaClass):
|
||||
@staticmethod
|
||||
def custom_function(name:str, *src:UOp) -> UOp: return UOp(Ops.CUSTOM_FUNCTION, src=src, arg=name)
|
||||
|
||||
# opaque bodies are just CALLs; value-producing bodies become CALLs with RETURNED placeholders as extra inputs
|
||||
# opaque bodies are just CALLs; value-producing bodies become CALLs with unbound BUFFER placeholders as extra inputs
|
||||
_OPAQUE_CALL_BODIES = {Ops.SINK, Ops.PROGRAM, Ops.LINEAR, Ops.COPY, Ops.CUSTOM_FUNCTION}
|
||||
def call(self, *srcs:UOp, ret_dtype:DType|None=None, grad_fxn:Callable|None=None,
|
||||
name:str|None=None, precompile:bool=False, precompile_backward:bool=False, aux:Any=None) -> UOp:
|
||||
if ret_dtype is not None: return UOp(Ops.CALL, src=(self,)+srcs, arg=ret_dtype)
|
||||
# calls are launched per device, so an open DEVICE range is allowed to cross the call boundary
|
||||
assert all(r.arg[-1] is AxisType.DEVICE for r in self.ranges), \
|
||||
f"ranges {self.ranges} are leaking out of the call in {self.pyrender()}"
|
||||
if self.op in UOp._OPAQUE_CALL_BODIES:
|
||||
return UOp(Ops.CALL, src=(self,)+srcs, arg=CallInfo(grad_fxn, name, precompile, precompile_backward, aux))
|
||||
# the (possibly void) return dtype lives in the CallInfo; an external C call is a CALL on a CUSTOM_FUNCTION
|
||||
# body holding the callee (a function pointer), rendered as an indirect call
|
||||
return UOp(Ops.CALL, src=(self,)+srcs, arg=CallInfo(grad_fxn, name, precompile, precompile_backward, aux,
|
||||
ret_dtype if ret_dtype is not None else dtypes.void))
|
||||
assert ret_dtype is None, "ret_dtype requires an opaque body, use a CUSTOM_FUNCTION body for external calls"
|
||||
# value-producing bodies delegate to call_outputs with a single output
|
||||
return UOp.call_outputs((self,), *srcs, grad_fxn=grad_fxn, name=name, precompile=precompile,
|
||||
precompile_backward=precompile_backward, aux=aux)
|
||||
@@ -1302,11 +1305,13 @@ class CallInfo:
|
||||
precompile: bool = False
|
||||
precompile_backward: bool = False
|
||||
aux: Any = None
|
||||
dtype: DType = dtypes.void
|
||||
# grad_fxn can't be pickled
|
||||
def __reduce__(self): return (CallInfo, (None, self.name, self.precompile, self.precompile_backward, self.aux))
|
||||
def __reduce__(self): return (CallInfo, (None, self.name, self.precompile, self.precompile_backward, self.aux, self.dtype))
|
||||
def __repr__(self):
|
||||
gf = id(self.grad_fxn) if self.grad_fxn else None
|
||||
return f"CallInfo({gf}, {repr(self.name)}, {self.precompile}, {self.precompile_backward})"
|
||||
return f"CallInfo({gf}, {repr(self.name)}, {self.precompile}, {self.precompile_backward})" + \
|
||||
(f", {self.dtype}" if self.dtype is not dtypes.void else "")
|
||||
|
||||
# ******** ops in python ********
|
||||
|
||||
@@ -1694,10 +1699,8 @@ class RewriteContext:
|
||||
continue
|
||||
# no rewrite, process children then come back to rebuild
|
||||
stack.append((n, True))
|
||||
# program bodies (kernels, value calls) are never rewritten separately unless the rewrite explicitly enters
|
||||
# calls; other call graphs (dtype-arg calls) are plain dataflow and always rewritten
|
||||
if n.op is Ops.CALL and not self.enter_calls and n.src[0].op in UOp._OPAQUE_CALL_BODIES:
|
||||
self.replace[n.src[0]] = n.src[0]
|
||||
# CALL bodies are never rewritten separately, rewrites that need them pass enter_calls=True
|
||||
if n.op is Ops.CALL and not self.enter_calls: self.replace[n.src[0]] = n.src[0]
|
||||
for x in reversed(n.src):
|
||||
if x not in self.replace: stack.append((x, False))
|
||||
else:
|
||||
@@ -1734,10 +1737,9 @@ class RewriteContext:
|
||||
if n in waitlist: stack.extend(waitlist.pop(n))
|
||||
continue
|
||||
stack.append((n, 1, new_n))
|
||||
# NOTE: CALLs are handled as a special case: program bodies are not included in the graph_rewrite unless the
|
||||
# rewrite explicitly enters calls (a CALL of an address is not a body, its srcs are regular dataflow)
|
||||
if new_n.op is Ops.CALL and not self.enter_calls and new_n.src[0].op in UOp._OPAQUE_CALL_BODIES:
|
||||
self.replace[new_n.src[0]] = new_n.src[0]
|
||||
# NOTE: CALLs are handled as a special case: their bodies are not included in the graph_rewrite,
|
||||
# rewrites that need them pass enter_calls=True
|
||||
if new_n.op is Ops.CALL and not self.enter_calls: self.replace[new_n.src[0]] = new_n.src[0]
|
||||
for x in reversed(new_n.src):
|
||||
if x in on_stack: continue
|
||||
stack.append((x, 0, x))
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import math, functools
|
||||
from typing import Any
|
||||
from tinygrad.uop.ops import PatternMatcher, UPat, GroupOp, Ops, UOp, AxisType, KernelInfo, ParamArg
|
||||
from tinygrad.uop.ops import PatternMatcher, UPat, GroupOp, Ops, UOp, AxisType, KernelInfo, ParamArg, CallInfo
|
||||
from tinygrad.uop.render import print_uops, pyrender
|
||||
from tinygrad.dtype import DType, dtypes, AddrSpace, Invalid, ConstFloat
|
||||
from tinygrad.helpers import DEBUG, Context, SPEC, Metadata, panic, CHECK_OOB, all_same, is_image_shape
|
||||
@@ -102,9 +102,11 @@ spec_shared = PatternMatcher([
|
||||
(UPat((Ops.CUSTOMI, Ops.CUSTOM), name="x"),
|
||||
lambda x: isinstance(x.arg, tuple) and len(x.arg) == 2 and isinstance(x.arg[0], str) and isinstance(x.arg[1], DType)),
|
||||
|
||||
# CALL of an external function
|
||||
(UPat(Ops.CALL, src=(UPat(),), allow_any_len=True, name="x"),
|
||||
lambda x: matches_dtype(x.src[0], dtypes.uint64) and isinstance(x.arg, DType) if x.src[0].dtype is not dtypes.void else None),
|
||||
# a CUSTOM_FUNCTION with srcs is the body of an external call, holding the callee (a function pointer)
|
||||
(UPat(Ops.CUSTOM_FUNCTION, name="x", allow_any_len=True), lambda x: isinstance(x.arg, str)),
|
||||
# CALL: the body is always an opaque body, the arg is a CallInfo stating the (possibly void) dtype
|
||||
(UPat(Ops.CALL, src=(UPat(tuple(UOp._OPAQUE_CALL_BODIES)),), allow_any_len=True, name="x"),
|
||||
lambda x: isinstance(x.arg, CallInfo) and x.dtype is x.arg.dtype),
|
||||
|
||||
# pattern compiler IR ops (not in tensor/program graphs, but spec-compliant)
|
||||
(UPat(Ops.PYLITERAL), lambda: True),
|
||||
@@ -149,9 +151,6 @@ spec_tensor = PatternMatcher([
|
||||
# custom function
|
||||
(UPat(Ops.CUSTOM_FUNCTION, name="x"), lambda x: isinstance(x.arg, str)),
|
||||
|
||||
# CALL
|
||||
(UPat(Ops.CALL, dtypes.void, src=(UPat((Ops.SINK, Ops.LINEAR, Ops.PROGRAM, Ops.COPY, Ops.CUSTOM_FUNCTION)),), allow_any_len=True), lambda: True),
|
||||
|
||||
# SPECIAL is index before index lowering. custom_kernel currently has this
|
||||
(UPat(Ops.SPECIAL, src=(UPat(dtype=dtypes.weakint),), name="s"), lambda s: isinstance(s.arg, str)),
|
||||
|
||||
@@ -260,8 +259,8 @@ spec_kernel_graph = PatternMatcher([
|
||||
# mstack/mselect
|
||||
(UPat(Ops.MSTACK, name="x"), lambda x: all(isinstance(s.device, str) for s in x.src) or (all_same(x.src) and x.src[0].device is None)),
|
||||
(UPat(Ops.MSELECT, name="x"), lambda x: isinstance(x.src[0].device, tuple) and x.arg < len(x.src[0].device)),
|
||||
# all calls are on various sinks
|
||||
(UPat(Ops.CALL, src=(UPat((Ops.SINK, Ops.LINEAR, Ops.PROGRAM, Ops.CUSTOM_FUNCTION)),), allow_any_len=True), lambda: True),
|
||||
# all calls are on opaque bodies
|
||||
(UPat(Ops.CALL, src=(UPat(tuple(UOp._OPAQUE_CALL_BODIES)),), allow_any_len=True), lambda: True),
|
||||
# after on PARAM or AFTER
|
||||
(UPat(Ops.AFTER, src=(UPat(GroupOp.Movement.union({Ops.PARAM, Ops.AFTER, Ops.BUFFER, Ops.MSTACK, Ops.MSELECT, Ops.BITCAST, Ops.RESHAPE})),),
|
||||
allow_any_len=True), lambda: True),
|
||||
|
||||
Reference in New Issue
Block a user