restrict allowed call bodies (#17925)
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) Canceled after 0s
Benchmarks / Mac pytest (push) Canceled after 0s
Benchmarks / LLM (DEV=AMD) (push) Canceled after 0s
Benchmarks / LLM (DEV=METAL) (push) Canceled after 0s
Benchmarks / LLM (DEV=NV) (push) Canceled after 0s
Benchmarks / HLB-CIFAR10 (DEV=AMD) (push) Canceled after 0s
Benchmarks / HLB-CIFAR10 (DEV=METAL) (push) Canceled after 0s
Benchmarks / HLB-CIFAR10 (DEV=NV) (push) Canceled after 0s
Benchmarks / MLPerf (AMD) (push) Canceled after 0s
Benchmarks / MLPerf (NV) (push) Canceled after 0s
Benchmarks / Stable Diffusion (DEV=AMD) (push) Canceled after 0s
Benchmarks / Stable Diffusion (DEV=METAL) (push) Canceled after 0s
Benchmarks / Stable Diffusion (DEV=NV) (push) Canceled after 0s
Benchmarks / Multi-GPU Benchmarks (DEV=AMD) (push) Canceled after 0s
Benchmarks / Multi-GPU Benchmarks (DEV=NV) (push) Canceled after 0s
Benchmarks / Tests (DEV=AMD) (push) Canceled after 0s
Benchmarks / Tests (DEV=METAL) (push) Canceled after 0s
Benchmarks / Tests (DEV=NV) (push) Canceled after 0s
Benchmarks / UsbGPU Benchmark (push) Canceled after 0s
Benchmarks / openpilot 0.11.2 compile3 dmonitoring (DEV=QCOM) (push) Canceled after 0s
Benchmarks / openpilot 0.11.0 compile3 dmonitoring (DEV=QCOM) (push) Canceled after 0s
Benchmarks / openpilot 0.11.0 compile3 policy (DEV=QCOM) (push) Canceled after 0s
Benchmarks / openpilot 0.11.2 compile3 supercombo (DEV=QCOM) (push) Canceled after 0s
Benchmarks / openpilot 0.11.0 compile3 vision (DEV=QCOM) (push) Canceled after 0s
Benchmarks / openpilot 0.11.0 compile3 dmonitoring (DEV=QCOM:IR3) (push) Canceled after 0s
Benchmarks / openpilot 0.11.2 compile3 dmonitoring (DEV=QCOM:IR3) (push) Canceled after 0s
Benchmarks / openpilot 0.11.0 compile3 policy (DEV=QCOM:IR3) (push) Canceled after 0s
Benchmarks / openpilot 0.11.2 compile3 supercombo (DEV=QCOM:IR3) (push) Canceled after 0s
Benchmarks / openpilot 0.11.0 compile3 vision (DEV=QCOM:IR3) (push) Canceled after 0s
Benchmarks / DSP Benchmark (push) Canceled after 0s
Benchmarks / UsbGPU Benchmark (comma) (push) Canceled after 0s
Benchmarks / PCI Driver Benchmark (DEV=AMD) (push) Canceled after 0s
Benchmarks / PCI Driver Benchmark (DEV=NV) (push) Canceled after 0s
Benchmarks / LLVM Speed (push) Canceled after 0s
Platform Tests / MacOS (unit) (push) Canceled after 0s
Platform Tests / MacOS (unit, mock) (push) Canceled after 0s
Platform Tests / MacOS (DEV=METAL) (1) (push) Canceled after 0s
Platform Tests / MacOS (DEV=METAL) (2) (push) Canceled after 0s
Platform Tests / MacOS (DEV=CPU:CLANG) (push) Canceled after 0s
Platform Tests / MacOS (DEV=CPU:LLVM) (push) Canceled after 0s
Platform Tests / MacOS (DEV=CPU:LVP) (push) Canceled after 0s
Platform Tests / MacOS (DEV=WEBGPU) (push) Canceled after 0s
Platform Tests / Windows (DEV=CPU:CLANG) (push) Canceled after 0s
Platform Tests / Windows (DEV=CPU:LLVM) (push) Canceled after 0s
Platform Tests / Windows (DEV=CPU:X86) (push) Canceled after 0s
Platform Tests / Windows (DEV=WEBGPU) (push) Canceled after 0s

* restrict allowed call bodies

* simpler

* less

* simpler spec

* cleanup for spec
This commit is contained in:
George Hotz
2026-09-02 21:25:59 -07:00
committed by GitHub
parent 88face1a98
commit 6f4bfde234
8 changed files with 38 additions and 34 deletions
+3 -2
View File
@@ -5,12 +5,13 @@ from tinygrad.dtype import dtypes
from tinygrad.renderer.cstyle import CStyleLanguage from tinygrad.renderer.cstyle import CStyleLanguage
from tinygrad.uop.ops import KernelInfo 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: 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")) 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: 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")) 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") @unittest.skipUnless(isinstance(Device["CPU"].renderer, CStyleLanguage), "TODO: CALL is rendered in C style only")
+1 -1
View File
@@ -24,7 +24,7 @@ def _make_linear(buffer_lists, copies=None):
src0 = bufs[0].copy_to_device(bufs[1].device) src0 = bufs[0].copy_to_device(bufs[1].device)
else: else:
src0 = UOp(Ops.SINK, src=tuple(bufs)) 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)) return UOp(Ops.LINEAR, src=tuple(calls))
def _get_planned_view(buf:UOp) -> tuple[UOp, int, int]|None: def _get_planned_view(buf:UOp) -> tuple[UOp, int, int]|None:
+1 -1
View File
@@ -228,7 +228,7 @@ class TestViz(unittest.TestCase):
pm = PatternMatcher([(UPat(Ops.CONST, arg=3, name="x"), lambda x: UOp.const(4, x.dtype))]) pm = PatternMatcher([(UPat(Ops.CONST, arg=3, name="x"), lambda x: UOp.const(4, x.dtype))])
with save_viz() as viz: with save_viz() as viz:
inner = UOp.const(3) 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) graph_rewrite(call, TrackedPatternMatcher(pm.patterns), enter_calls=True)
details = list(viz.get_details(0, 0)) details = list(viz.get_details(0, 0))
self.assertTrue(details[-1]["change"], "viz replay should detect change inside CALL") self.assertTrue(details[-1]["change"], "viz replay should detect change inside CALL")
+1 -1
View File
@@ -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 # 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: 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 # 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 # final symbolic before decomp
sink = graph_rewrite(sink, symbolic, name="final symbolic") sink = graph_rewrite(sink, symbolic, name="final symbolic")
+4 -4
View File
@@ -65,9 +65,9 @@ base_rewrite = PatternMatcher([
(UPat(GroupOp.ALU, name="x"), lambda ctx,x: ctx.code_for_op[x.op]( (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)), *([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 # 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(),), allow_any_len=True, name="x"), lambda ctx,x: (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[x.src[0]]}))" + 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 "")), 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 # custom passes through with format
@@ -214,7 +214,7 @@ class CStyleLanguage(Renderer):
c: defaultdict[str, int] = defaultdict(int) c: defaultdict[str, int] = defaultdict(int)
name = "test" name = "test"
for u in uops: 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 == Ops.STACK and len(u.src) == 0: continue
if u.op is Ops.AFTER: if u.op is Ops.AFTER:
r[u] = r[u.src[0]] r[u] = r[u.src[0]]
+3 -1
View File
@@ -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 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: 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) *[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 def usb_bulk(devs, dep, endpoint:int, data:UOp, length, timeout:int=1000) -> UOp: # NULL actual_length out param
+17 -15
View File
@@ -126,8 +126,8 @@ def dtype_from_uop(op:Ops, src:tuple[UOp,...], arg:Any) -> DType:
# always void # always void
return dtypes.void return dtypes.void
case Ops.CALL: 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 # a call states its (possibly void) dtype in the CallInfo
return arg if isinstance(arg, DType) else dtypes.void return arg.dtype if isinstance(arg, CallInfo) else dtypes.void
case Ops.CUSTOM | Ops.CUSTOMI: 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}" 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] return arg[1]
@@ -1202,16 +1202,19 @@ class UOp(RandMixin, metaclass=UOpMetaClass):
@staticmethod @staticmethod
def custom_function(name:str, *src:UOp) -> UOp: return UOp(Ops.CUSTOM_FUNCTION, src=src, arg=name) 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} _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, 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: 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 # 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), \ 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()}" f"ranges {self.ranges} are leaking out of the call in {self.pyrender()}"
if self.op in UOp._OPAQUE_CALL_BODIES: 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 # 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, return UOp.call_outputs((self,), *srcs, grad_fxn=grad_fxn, name=name, precompile=precompile,
precompile_backward=precompile_backward, aux=aux) precompile_backward=precompile_backward, aux=aux)
@@ -1302,11 +1305,13 @@ class CallInfo:
precompile: bool = False precompile: bool = False
precompile_backward: bool = False precompile_backward: bool = False
aux: Any = None aux: Any = None
dtype: DType = dtypes.void
# grad_fxn can't be pickled # 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): def __repr__(self):
gf = id(self.grad_fxn) if self.grad_fxn else None 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 ******** # ******** ops in python ********
@@ -1694,10 +1699,8 @@ class RewriteContext:
continue continue
# no rewrite, process children then come back to rebuild # no rewrite, process children then come back to rebuild
stack.append((n, True)) stack.append((n, True))
# program bodies (kernels, value calls) are never rewritten separately unless the rewrite explicitly enters # CALL bodies are never rewritten separately, rewrites that need them pass enter_calls=True
# calls; other call graphs (dtype-arg calls) are plain dataflow and always rewritten if n.op is Ops.CALL and not self.enter_calls: self.replace[n.src[0]] = n.src[0]
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]
for x in reversed(n.src): for x in reversed(n.src):
if x not in self.replace: stack.append((x, False)) if x not in self.replace: stack.append((x, False))
else: else:
@@ -1734,10 +1737,9 @@ class RewriteContext:
if n in waitlist: stack.extend(waitlist.pop(n)) if n in waitlist: stack.extend(waitlist.pop(n))
continue continue
stack.append((n, 1, new_n)) 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 # NOTE: CALLs are handled as a special case: their bodies are not included in the graph_rewrite,
# rewrite explicitly enters calls (a CALL of an address is not a body, its srcs are regular dataflow) # rewrites that need them pass enter_calls=True
if new_n.op is Ops.CALL and not self.enter_calls and new_n.src[0].op in UOp._OPAQUE_CALL_BODIES: if new_n.op is Ops.CALL and not self.enter_calls: self.replace[new_n.src[0]] = new_n.src[0]
self.replace[new_n.src[0]] = new_n.src[0]
for x in reversed(new_n.src): for x in reversed(new_n.src):
if x in on_stack: continue if x in on_stack: continue
stack.append((x, 0, x)) stack.append((x, 0, x))
+8 -9
View File
@@ -1,6 +1,6 @@
import math, functools import math, functools
from typing import Any 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.uop.render import print_uops, pyrender
from tinygrad.dtype import DType, dtypes, AddrSpace, Invalid, ConstFloat 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 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"), (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)), 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 # a CUSTOM_FUNCTION with srcs is the body of an external call, holding the callee (a function pointer)
(UPat(Ops.CALL, src=(UPat(),), allow_any_len=True, name="x"), (UPat(Ops.CUSTOM_FUNCTION, name="x", allow_any_len=True), lambda x: isinstance(x.arg, str)),
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), # 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) # pattern compiler IR ops (not in tensor/program graphs, but spec-compliant)
(UPat(Ops.PYLITERAL), lambda: True), (UPat(Ops.PYLITERAL), lambda: True),
@@ -149,9 +151,6 @@ spec_tensor = PatternMatcher([
# custom function # custom function
(UPat(Ops.CUSTOM_FUNCTION, name="x"), lambda x: isinstance(x.arg, str)), (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 # 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)), (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 # 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.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)), (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 # all calls are on opaque bodies
(UPat(Ops.CALL, src=(UPat((Ops.SINK, Ops.LINEAR, Ops.PROGRAM, Ops.CUSTOM_FUNCTION)),), allow_any_len=True), lambda: True), (UPat(Ops.CALL, src=(UPat(tuple(UOp._OPAQUE_CALL_BODIES)),), allow_any_len=True), lambda: True),
# after on PARAM or AFTER # 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})),), (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), allow_any_len=True), lambda: True),