forked from tinygrad/tinygrad
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0f7226da30 | ||
|
|
ab3a217c0b |
@@ -59,6 +59,7 @@ class TestLinearizer(unittest.TestCase):
|
||||
self.assertEqual(len([x for x in uops if x.op is Ops.CAST]), 0)
|
||||
|
||||
@unittest.skipIf(isinstance(Device[Device.DEFAULT].renderer, PTXRenderer), "broken on ptx")
|
||||
@unittest.skipUnless(Device[Device.DEFAULT].renderer.supports_ranges, "test inspects ranges, which are rewritten to loops on this renderer")
|
||||
def test_late_bias_load(self):
|
||||
img = Tensor.empty(1, 3, 16, 16)
|
||||
w = Tensor.empty(16, 3, 3, 3)
|
||||
@@ -238,6 +239,7 @@ class TestLinearizer(unittest.TestCase):
|
||||
helper_arg_acc_dtype(d.conv2d(w, dtype=acc_dtype), expected_dtype)
|
||||
|
||||
@unittest.skipUnless(Device[Device.DEFAULT].renderer.supports_float4, "test requires float4")
|
||||
@unittest.skipUnless(Device[Device.DEFAULT].renderer.supports_ranges, "test inspects ranges, which are rewritten to loops on this renderer")
|
||||
def test_simple_unroll_no_between_phi_dependencies(self):
|
||||
x, y = Tensor.empty(64, 64), Tensor.empty(64, 64)
|
||||
r = (x@y).relu()
|
||||
@@ -299,6 +301,7 @@ class TestLinearizer(unittest.TestCase):
|
||||
program = to_program(replace_opts(linear.src[-1].src[0], []), renderer=Device[Device.DEFAULT].renderer)
|
||||
assert not any(u.op == Ops.WHERE for u in tuple(program.src[1].src)), "found where where where should be folded"
|
||||
|
||||
@unittest.skipUnless(Device[Device.DEFAULT].renderer.supports_ranges, "test inspects ranges, which are rewritten to loops on this renderer")
|
||||
def test_phi_simplification(self):
|
||||
def helper(t, max_ops=0):
|
||||
ast = helper_linearizer_opt(t)
|
||||
|
||||
@@ -22,7 +22,7 @@ from tinygrad.codegen.opt.postrange import apply_opts
|
||||
from tinygrad.codegen.late.gater import pm_move_gates_from_index
|
||||
from tinygrad.codegen.simplify import pm_simplify_ranges, pm_flatten_range, pm_split_ranges, pm_load_collapse
|
||||
from tinygrad.schedule.rangeify import pm_mops
|
||||
from tinygrad.codegen.late.linearizer import CFGContext, pm_split_ends, pm_add_control_flow, linearize
|
||||
from tinygrad.codegen.late.linearizer import CFGContext, pm_split_ends, pm_add_control_flow, linearize, ranges_to_loops
|
||||
from tinygrad.codegen.late.regalloc import LinearScanRegallocContext, pm_regalloc_rewrite
|
||||
from tinygrad.codegen.late.coalesce import memory_coalescing, pm_simplify_add_image
|
||||
from tinygrad.helpers import all_same, flatten, argsort, partition
|
||||
@@ -349,6 +349,12 @@ def full_rewrite_to_sink(ast:UOp, ren:Renderer, optimize:bool=True) -> UOp:
|
||||
pm_final_rewrite = pm_decomp+extra_matcher+pm_split_ends
|
||||
sink = graph_rewrite(sink, pm_final_rewrite+pm_remove_invalid, ctx=ren, name="final rewrite")
|
||||
|
||||
if VIZ: graph_rewrite(sink, PatternMatcher([]), name="View Output AST")
|
||||
if SPEC: type_verify(sink, spec_program)
|
||||
|
||||
# rewrite bounded ranges to loops for renderers without range support, after validation like instruction selection
|
||||
if not ren.supports_ranges: sink = ranges_to_loops(sink)
|
||||
|
||||
# this was the linearizer
|
||||
sink = graph_rewrite(sink, pm_add_control_flow, ctx=CFGContext(sink), name="add control flow", bottom_up=True)
|
||||
|
||||
@@ -356,9 +362,6 @@ def full_rewrite_to_sink(ast:UOp, ren:Renderer, optimize:bool=True) -> UOp:
|
||||
num_params = len([x for x in sink.toposort() if x.op is Ops.PARAM and x.arg.slot != -1])
|
||||
sink = graph_rewrite(sink, pm_number_params, ctx=[num_params], name="number params with -1", walk=True)
|
||||
|
||||
if VIZ: graph_rewrite(sink, PatternMatcher([]), name="View Output AST")
|
||||
if SPEC: type_verify(sink, spec_program)
|
||||
|
||||
# return the rewritten sink
|
||||
return sink
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import heapq
|
||||
from typing import Any
|
||||
from collections import defaultdict
|
||||
from tinygrad.uop.ops import PatternMatcher, UOp, Ops, UPat, multirange_str
|
||||
from tinygrad.uop.ops import PatternMatcher, UOp, Ops, UPat, multirange_str, ParamArg, AxisType
|
||||
from tinygrad.dtype import AddrSpace, dtypes
|
||||
from tinygrad.helpers import prod, getenv, TUPLE_ORDER
|
||||
|
||||
@@ -93,3 +93,42 @@ pm_split_ends = PatternMatcher([
|
||||
# split the ends
|
||||
(UPat(Ops.END, name="e"), do_split_ends),
|
||||
])
|
||||
|
||||
def ranges_to_loops(sink:UOp) -> UOp:
|
||||
# rewrite bounded ranges to bound-less loops with a register counter: i = 0; loop { body; i += 1; loop again while i < bound }
|
||||
slot = max((u.arg.slot for u in sink.toposort() if u.op is Ops.BUFFER and u.addrspace == AddrSpace.REG), default=-1) + 1
|
||||
ends = [u for u in sink.toposort() if u.op is Ops.END and any(x.op is Ops.RANGE and x.dtype is not dtypes.void for x in u.src[1:])]
|
||||
# e.ranges over-approximates nesting (it flows ranges through ordering deps), so compute true nesting from the body slices
|
||||
# NOTE: uop identity is not stable (the uop cache is weak), all lookups are by uop key
|
||||
end_for_range = {r.key: e for e in ends for r in e.src[1:] if r.op is Ops.RANGE and r.dtype is not dtypes.void}
|
||||
body_ends = {e.key: {u.key for u in e.src[0].toposort()} for e in ends}
|
||||
repl: dict[UOp, UOp] = {}
|
||||
range_to_loop: dict[bytes, UOp] = {}
|
||||
for e in ends:
|
||||
# the counter init is placed after the enclosing loops so it resets every outer iteration, the loop header depends on it so it runs first
|
||||
enclosing = tuple(r for r in e.ranges if (er:=end_for_range.get(r.key)) is not None and e.key in body_ends[er.key])
|
||||
e = e.substitute(repl)
|
||||
assert len(e.src) == 2, f"expected a split END with one range, got {len(e.src)-1} ranges"
|
||||
r = e.src[1]
|
||||
i = UOp(Ops.BUFFER, src=(UOp.const(dtypes.int, 1),), arg=ParamArg(slot, r.dtype, addrspace=AddrSpace.REG))
|
||||
slot += 1
|
||||
z = UOp.const(dtypes.int, 0)
|
||||
init = i.after(*enclosing).index(z).store(UOp.const(r.dtype, 0))
|
||||
i = i.after(init)
|
||||
# a do-while can't skip its first iteration, so a range with a possibly zero bound gets a one-time entry guard on the loop header
|
||||
guard = () if r.src[0].vmin >= 1 else (UOp.const(r.dtype, 0) < r.src[0],)
|
||||
l = range_to_loop[r.key] = UOp(Ops.RANGE, dtypes.void, src=(init,)+guard, arg=(r.arg[0], AxisType.LOOP))
|
||||
iv = i.after(l).index(z).load()
|
||||
inc = iv + UOp.const(r.dtype, 1)
|
||||
body = e.src[0].substitute({r: iv})
|
||||
# the counter store is part of the loop body, an AFTER body can't be in a GROUP so sequence it with a dep instead
|
||||
ret = body.after(i.index(z).store(inc)) if body.op is Ops.AFTER else UOp.group(body, i.index(z).store(inc))
|
||||
repl[e] = ret.end(l, inc < r.src[0])
|
||||
# keep the tracked loop headers up to date: their init deps on enclosing ranges get rewritten by the same substitution
|
||||
for k in range_to_loop: range_to_loop[k] = range_to_loop[k].substitute({r: iv})
|
||||
if not len(repl): return sink
|
||||
out = sink.substitute(repl)
|
||||
# ordering deps on the old ranges (scope AFTERs outside the loop bodies) point at the loop headers
|
||||
fix = {a: a.replace(src=(a.src[0],) + tuple(range_to_loop[s.key] if s.key in range_to_loop else s for s in a.src[1:]))
|
||||
for a in out.toposort() if a.op is Ops.AFTER and any(s.key in range_to_loop for s in a.src[1:])}
|
||||
return out.substitute(fix) if len(fix) else out
|
||||
|
||||
@@ -73,6 +73,8 @@ class Renderer:
|
||||
tensor_cores: list[TensorCore] = []
|
||||
extra_matcher: PatternMatcher|None = None
|
||||
code_for_op: dict[Ops, Callable] = {}
|
||||
# renderers without range support get all bounded ranges rewritten to loops in codegen
|
||||
supports_ranges: bool = True
|
||||
|
||||
compiler: Compiler = Compiler()
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@ from tinygrad.codegen.opt import tc
|
||||
from tinygrad.renderer import Renderer
|
||||
from tinygrad.renderer.cstyle import HIPRenderer, create_non_native_float_pats, pm_manual_bf16_cast
|
||||
from tinygrad.codegen.decomp.transcendental import xexp2, xlog2
|
||||
from tinygrad.uop.ops import UOp, PatternMatcher, UPat, Ops, GroupOp, range_str
|
||||
from tinygrad.uop.ops import UOp, PatternMatcher, UPat, Ops, GroupOp
|
||||
from tinygrad.dtype import dtypes, float_to_fp8, DType, truncate, AddrSpace
|
||||
from tinygrad.helpers import prod, Target, CPU_COUNT, getenv, OSX
|
||||
|
||||
@@ -101,28 +101,13 @@ base_rewrite = PatternMatcher([
|
||||
(UPat(Ops.WHERE, name="x"), lambda ctx,x:
|
||||
f" {ctx[x]} = select {ldt(x.src[0].dtype)} {ctx[x.src[0]]}, {ldt(x.src[1].dtype)} {ctx[x.src[1]]}, {ldt(x.src[2].dtype)} {ctx[x.src[2]]}"),
|
||||
|
||||
# loop (a RANGE with no src is an unbounded loop header)
|
||||
(UPat(Ops.RANGE, dtypes.void, name="l"), lambda ctx,l: f" br label %loop_{ctx[l][1:]}\nloop_{ctx[l][1:]}:"),
|
||||
# loop (ranges are rewritten to loops in codegen), a bool src is a one-time entry guard for possibly zero trip counts
|
||||
(UPat(Ops.RANGE, dtypes.void, name="l"), lambda ctx,l:
|
||||
f" br i1 {ctx[g]}, label %loop_{ctx[l][1:]}, label %loop_exit_{ctx[l][1:]}\nloop_{ctx[l][1:]}:" \
|
||||
if (g:=next((s for s in l.src if s.dtype is dtypes.bool), None)) is not None else f" br label %loop_{ctx[l][1:]}\nloop_{ctx[l][1:]}:"),
|
||||
(UPat(Ops.END, src=(UPat(), UPat(Ops.RANGE, dtypes.void, name="l"), UPat(name="c"))), lambda ctx,l,c:
|
||||
f" br i1 {ctx[c]}, label %loop_{ctx[l][1:]}, label %loop_exit_{ctx[l][1:]}\nloop_exit_{ctx[l][1:]}:"),
|
||||
|
||||
# range
|
||||
(UPat(Ops.RANGE, name="r"), lambda ctx,r:
|
||||
f" br label %loop_entry_{range_str(r)}\n"
|
||||
f"loop_entry_{range_str(r)}:\n"
|
||||
f" br label %loop_latch_{range_str(r)}\n"
|
||||
f"loop_latch_{range_str(r)}:\n"
|
||||
f" {ctx[r]} = phi {ldt(r.dtype)} [ 0, %loop_entry_{range_str(r)} ], [ {ctx[r]}phi, %loop_footer_{range_str(r)} ]\n"
|
||||
f" {ctx[r]}phi = add {ldt(r.dtype)} {ctx[r]}, 1\n"
|
||||
f" {ctx[r]}cmp = icmp ult {ldt(r.dtype)} {ctx[r]}, {ctx[r.src[0]]}\n"
|
||||
f" br i1 {ctx[r]}cmp, label %loop_body_{range_str(r)}, label %loop_exit_{range_str(r)}\n"
|
||||
f"loop_body_{range_str(r)}:"),
|
||||
(UPat(Ops.END, src=(UPat(), UPat(Ops.RANGE, name="r"))), lambda r:
|
||||
f" br label %loop_footer_{range_str(r)}\n"
|
||||
f"loop_footer_{range_str(r)}:\n"
|
||||
f" br label %loop_latch_{range_str(r)}\n"
|
||||
f"loop_exit_{range_str(r)}:"),
|
||||
|
||||
# if
|
||||
(UPat(Ops.IF, name="x"), lambda ctx,x: f" br i1 {ctx[x.src[0]]}, label %ifbody_{ctx[x][1:]}, label %ifskip_{ctx[x][1:]}\nifbody_{ctx[x][1:]}:"),
|
||||
(UPat(Ops.ENDIF, name="x"), lambda ctx,x: f" br label %ifskip_{ctx[x.src[0]][1:]}\nifskip_{ctx[x.src[0]][1:]}:"),
|
||||
@@ -132,6 +117,7 @@ base_rewrite = PatternMatcher([
|
||||
|
||||
class LLVMRenderer(Renderer):
|
||||
supports_float4 = True
|
||||
supports_ranges = False
|
||||
abi: str | None
|
||||
string_rewrite: PatternMatcher
|
||||
code_for_op = {k:lambda:None for v in lop.values() for k in v.keys()}
|
||||
|
||||
+13
-25
@@ -3,7 +3,7 @@ from tinygrad.dtype import AddrSpace, DType, dtypes, truncate
|
||||
from tinygrad.helpers import DEBUG, OSX, unwrap, fromimport, Target, is_image_shape
|
||||
from tinygrad.renderer import Renderer
|
||||
from tinygrad.renderer.cstyle import CUDARenderer, OpenCLRenderer
|
||||
from tinygrad.uop.ops import GroupOp, Ops, UOp, PatternMatcher, UPat, range_str
|
||||
from tinygrad.uop.ops import GroupOp, Ops, UOp, PatternMatcher, UPat
|
||||
from tinygrad.runtime.autogen import mesa, libc
|
||||
from tinygrad.runtime.support.c import POINTER
|
||||
import base64, ctypes, struct, functools, inspect, itertools
|
||||
@@ -117,6 +117,7 @@ def nidx(b:mesa.nir_builder, buf, off, space, itemsize, gate=None) -> mesa.nir_d
|
||||
class NIRRenderer(Renderer):
|
||||
suffix = "NIR"
|
||||
nir_options: bytes
|
||||
supports_ranges = False
|
||||
global_max, local_max, shared_max = CUDARenderer.global_max, CUDARenderer.local_max, CUDARenderer.shared_max
|
||||
code_for_op = {**{k:lambda:None for k in u_aop.keys()}, **{k:lambda:None for k in s_aop.keys()}, **{k:lambda:None for k in f_aop.keys()}}
|
||||
|
||||
@@ -187,8 +188,7 @@ class NIRRenderer(Renderer):
|
||||
self.prerender(uops)
|
||||
for u in [u for u in uops if u.op is Ops.SPECIAL and u.arg[0] == "l"]: self.b.shader.contents.info.workgroup_size[int(u.arg[-1])] = u.src[0].arg
|
||||
self.r: dict[UOp, Any] = {}
|
||||
self.param_idx = 0
|
||||
ranges: list[mesa.nir_def|None] = []
|
||||
self.param_idx, loop_ifs = 0, []
|
||||
|
||||
for u in uops:
|
||||
if u.op in {Ops.NOOP, Ops.GROUP} or (u.op is Ops.STACK and len(u.src) == 0): pass
|
||||
@@ -204,29 +204,17 @@ class NIRRenderer(Renderer):
|
||||
self.r[u] = nimm(self.b, self.b.shader.contents.info.shared_size, dtypes.long)
|
||||
self.b.shader.contents.info.shared_size += u.max_numel()*u.dtype.itemsize
|
||||
elif u.op == Ops.RANGE:
|
||||
if u.dtype == dtypes.void:
|
||||
# a RANGE with no bound is a loop header: just open the loop, the END adds the conditional backedge
|
||||
ranges.append(None)
|
||||
mesa.nir_push_loop(self.b)
|
||||
else:
|
||||
ranges.append(i:=deref_var(self.b, mesa.nir_local_variable_create(self.b.impl, glsl_type(u.dtype), f"idx{range_str(u)}".encode()).contents))
|
||||
nstore(self.b, AddrSpace.REG, i, nimm(self.b, 0, u.dtype))
|
||||
mesa.nir_push_loop(self.b)
|
||||
self.r[u] = nload(self.b, AddrSpace.REG, i, u)
|
||||
nif(self.b, nalu(self.b, "ilt", self.r[u], self.r[u.src[0]]), lambda: None, lambda: njump(self.b, mesa.nir_jump_break))
|
||||
# ranges are rewritten to loops in codegen: just open the loop, the END adds the conditional backedge
|
||||
# a bool src is a one-time entry guard for possibly zero trip counts
|
||||
assert u.dtype == dtypes.void, "NIRRenderer does not support ranges"
|
||||
guard = next((s for s in u.src if s.dtype is dtypes.bool), None)
|
||||
loop_ifs.append(mesa.nir_push_if(self.b, self.r[guard]) if guard is not None else None)
|
||||
mesa.nir_push_loop(self.b)
|
||||
elif u.op == Ops.END:
|
||||
r = u.src[1]
|
||||
if r.dtype == dtypes.void:
|
||||
# loop again while the condition is true
|
||||
nif(self.b, self.r[u.src[2]], lambda: None, lambda: njump(self.b, mesa.nir_jump_break))
|
||||
ranges.pop()
|
||||
mesa.nir_pop_loop(self.b, None)
|
||||
else:
|
||||
next_i = nalu(self.b, "iadd", self.r[r], nimm(self.b, 1, r.dtype))
|
||||
# TODO: this nif should be removable ... but TestMultiTensor.test_double_matmul_shard_W_0 segfaults with it gone
|
||||
nif(self.b, nalu(self.b, "ilt", next_i, self.r[r.src[0]]), lambda: None, lambda: njump(self.b, mesa.nir_jump_break))
|
||||
nstore(self.b, AddrSpace.REG, ranges.pop(), next_i),
|
||||
mesa.nir_pop_loop(self.b, None)
|
||||
# loop again while the condition is true
|
||||
nif(self.b, self.r[u.src[2]], lambda: None, lambda: njump(self.b, mesa.nir_jump_break))
|
||||
mesa.nir_pop_loop(self.b, None)
|
||||
if (nif_ref:=loop_ifs.pop()) is not None: mesa.nir_pop_if(self.b, nif_ref)
|
||||
else:
|
||||
d: mesa.nir_def|None = self.def_rewrite.rewrite(u, ctx=self)
|
||||
if d is None: raise RuntimeError(f"failed to render {u.op} srcs {[x.dtype for x in u.src]}")
|
||||
|
||||
@@ -116,18 +116,13 @@ string_rewrite = PatternMatcher([
|
||||
# simple
|
||||
(UPat(Ops.BUFFER, name="x"), lambda ctx, x: [] if x.addrspace == AddrSpace.REG else [
|
||||
f".shared .align 16 .b8 local{x.arg.slot}[{x.max_numel()*x.dtype.itemsize}];", f"mov.u64 {ctx.r[x]}, local{x.arg.slot}[0];"]),
|
||||
(UPat(Ops.RANGE, dtypes.void, name="l"), lambda ctx, l: f"WAITLOOP_{ctx.uops.index(l)}:"),
|
||||
# loop (ranges are rewritten to loops in codegen), a bool src is a one-time entry guard for possibly zero trip counts
|
||||
(UPat(Ops.RANGE, dtypes.void, name="l"), lambda ctx, l:
|
||||
[f"@!{ctx.r[g]} bra WAITLOOP_EXIT_{ctx.uops.index(l)};", f"WAITLOOP_{ctx.uops.index(l)}:"] \
|
||||
if (g:=next((s for s in l.src if s.dtype is dtypes.bool), None)) is not None else f"WAITLOOP_{ctx.uops.index(l)}:"),
|
||||
(UPat(Ops.END, src=(UPat(), UPat(Ops.RANGE, dtypes.void, name="l"), UPat(name="c"))), lambda ctx, l, c:
|
||||
f"@{ctx.r[c]} bra WAITLOOP_{ctx.uops.index(l)};"),
|
||||
(UPat(Ops.RANGE, name="r"), lambda ctx, r: [
|
||||
f"mov.u32 {ctx.r[r]}, -1;",
|
||||
f"bra END_{ctx.r[r][1:]};",
|
||||
"LOOP_" + f"{ctx.r[r][1:]}:"]),
|
||||
(UPat(Ops.END, name="x", src=(UPat(), UPat(Ops.RANGE, name="r"))), lambda ctx, x, r: [
|
||||
"END_" + f"{ctx.r[r][1:]}:",
|
||||
ctx.code_for_op[Ops.ADD](ctx.r[r], ctx.r[r], "1", dtypes.int, ctx.types[dtypes.int]),
|
||||
ctx.code_for_op[Ops.CMPLT](ctx.r[x], ctx.r[r], ctx.r[r.src[0]], dtypes.int, ctx.types[dtypes.int]),
|
||||
f"@{ctx.r[x]} bra LOOP_{ctx.r[r][1:]};"]),
|
||||
[f"@{ctx.r[c]} bra WAITLOOP_{ctx.uops.index(l)};"] +
|
||||
([f"WAITLOOP_EXIT_{ctx.uops.index(l)}:"] if any(s.dtype is dtypes.bool for s in l.src) else [])),
|
||||
(UPat(Ops.IF, name="x"), lambda ctx, x: f"@!{ctx.r[x.src[0]]} bra IF_{ctx.r[x.src[0]][1:]}_{ctx.uops.index(x)};"),
|
||||
(UPat(Ops.ENDIF, name="x"), lambda ctx, x: f"IF_{ctx.r[x.src[0].src[0]][1:]}_{ctx.uops.index(x.src[0])}:"),
|
||||
(UPat(Ops.WMMA, name="x"), lambda ctx, x: list(render_wmma(ctx, x))),
|
||||
@@ -136,6 +131,7 @@ string_rewrite = PatternMatcher([
|
||||
|
||||
class PTXRenderer(Renderer):
|
||||
suffix = "PTX"
|
||||
supports_ranges = False
|
||||
global_max, local_max, shared_max = CUDARenderer.global_max, CUDARenderer.local_max, CUDARenderer.shared_max
|
||||
tc_sm80 = [x for x in tc.cuda_sm80 if x.dtype_in in [dtypes.half, dtypes.float]]
|
||||
code_for_op = asm_for_op
|
||||
|
||||
@@ -475,6 +475,8 @@ class UOp(RandMixin, metaclass=UOpMetaClass):
|
||||
|
||||
@functools.cached_property
|
||||
def ended_ranges(self) -> tuple[UOp, ...]:
|
||||
# an END only ends ranges, the loop backedge condition is not an ended range
|
||||
if self.op is Ops.END: return tuple(x for x in self.src[1:] if x.op is Ops.RANGE)
|
||||
if self.op in range_start: return self.src[range_start[self.op]:]
|
||||
if self.op is Ops.AFTER: return tuple(flatten([x.ended_ranges for x in self.src[1:]]))
|
||||
return ()
|
||||
|
||||
Reference in New Issue
Block a user