forked from tinygrad/tinygrad
folding without UNMUL (#5628)
* folding without UNMUL * fix failures, index_collapse * import ReduceOps * test_arange_4096 isn't folding
This commit is contained in:
@@ -190,7 +190,6 @@ class TestLinearizer(unittest.TestCase):
|
||||
helper_linearizer_ast((store, ), [dataset, idxs], wanna_output=[real_index])
|
||||
|
||||
# AssertionError: repeated stores in uops
|
||||
@unittest.expectedFailure
|
||||
def test_argmax_multireduce_axis0(self):
|
||||
t = Tensor.randn(10, 20).realize()
|
||||
t_max = t.max((0,)).realize()
|
||||
|
||||
@@ -278,7 +278,7 @@ class TestLinearizerFailures(unittest.TestCase):
|
||||
LazyOp(op=BufferOps.CONST, src=(), arg=ConstBuffer(val=1.0, dtype=dtypes.float, st=ShapeTracker(views=(View(shape=(32640,), strides=(0,), offset=0, mask=None, contiguous=False),))))), arg=None),
|
||||
LazyOp(op=BufferOps.CONST, src=(), arg=ConstBuffer(val=0.0, dtype=dtypes.float, st=ShapeTracker(views=(View(shape=(32640,), strides=(0,), offset=0, mask=None, contiguous=False),))))), arg=None)), arg=None),), arg=(0,)),), arg=MemBuffer(idx=0, dtype=dtypes.float, st=ShapeTracker(views=(View(shape=(1,), strides=(0,), offset=0, mask=None, contiguous=True),)))),), arg=None)
|
||||
opts = [Opt(op=OptOps.GROUPTOP, axis=0, amt=16)]
|
||||
helper_test_lin(Kernel(ast), opts=opts, failed_platforms=["METAL", "GPU", "CUDA", "AMD", "NV"])
|
||||
helper_test_lin(Kernel(ast), opts=opts, failed_platforms=[])
|
||||
|
||||
# from fuzzing on metal
|
||||
def test_failure_34(self, unroll=False):
|
||||
@@ -337,7 +337,7 @@ class TestLinearizerFailures(unittest.TestCase):
|
||||
ast = LazyOp(op=MetaOps.KERNEL, src=(LazyOp(op=BufferOps.STORE, src=(LazyOp(op=BinaryOps.ADD, src=(LazyOp(op=ReduceOps.SUM, src=(LazyOp(op=BufferOps.CONST, src=(), arg=ConstBuffer(val=1, dtype=dtypes.int, st=ShapeTracker(views=(View(shape=(60001, 119999), strides=(0, 0), offset=0, mask=((0, 60001), (59999, 119999)), contiguous=False), View(shape=(60000, 60000), strides=(1, 120000), offset=0, mask=None, contiguous=False))))),), arg=(1,)), LazyOp(op=BufferOps.CONST, src=(), arg=ConstBuffer(val=-1, dtype=dtypes.int, st=ShapeTracker(views=(View(shape=(60000, 1), strides=(0, 0), offset=0, mask=None, contiguous=False),))))), arg=None),), arg=MemBuffer(idx=0, dtype=dtypes.int, st=ShapeTracker(views=(View(shape=(60000, 1), strides=(1, 0), offset=0, mask=None, contiguous=True),)))),), arg=None)
|
||||
for amt in [16,32]:
|
||||
opts = [Opt(op=OptOps.GROUPTOP, axis=0, amt=amt), Opt(op=OptOps.UNROLL, axis=0, amt=0)]
|
||||
helper_test_lin(Kernel(ast), opts=opts, failed_platforms=["METAL", "GPU"])
|
||||
helper_test_lin(Kernel(ast), opts=opts, failed_platforms=[])
|
||||
# END METAL=1 ./examples/beautiful_mnist.py failures
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
@@ -192,6 +192,9 @@ class TestOps(unittest.TestCase):
|
||||
def test_arange_big(self):
|
||||
helper_test_op([], lambda: torch.arange(256, dtype=torch.int32), lambda: Tensor.arange(256), forward_only=True)
|
||||
|
||||
def test_arange_4096(self):
|
||||
helper_test_op([], lambda: torch.arange(4096, dtype=torch.int32), lambda: Tensor.arange(4096), forward_only=True)
|
||||
|
||||
def test_sum_fake(self):
|
||||
helper_test_op([(256, 1)], lambda x: x.sum(axis=1))
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import unittest, itertools
|
||||
from test.helpers import TestUOps
|
||||
from tinygrad.dtype import dtypes
|
||||
from tinygrad.ops import BinaryOps, TernaryOps, UnaryOps # noqa: F401
|
||||
from tinygrad.ops import BinaryOps, TernaryOps, ReduceOps, UnaryOps # noqa: F401
|
||||
from tinygrad.codegen.uops import UOps, UOp, PatternMatcher, UPat, _match
|
||||
from tinygrad.codegen.uopgraph import UOpGraph, constant_folder
|
||||
|
||||
|
||||
@@ -117,14 +117,20 @@ def sum_collapse(phi_input, loop, val1, val2):
|
||||
return UOp(UOps.PHI, phi_input.dtype, (phi_input, v2))+ret
|
||||
return None
|
||||
|
||||
def loop_collapse(loop_start, loop_end, compval, idx, mval, multconst, rng):
|
||||
if getenv("DISABLE_LOOP_COLLAPSE") or not rng.arg[1]: return None # must be a REDUCE
|
||||
def loop_collapse(loop_start, loop_end, compval, idx, mval, multconst, rng, reduce_allow_any_len):
|
||||
if getenv("DISABLE_LOOP_COLLAPSE") or rng not in reduce_allow_any_len.src: return None # must be the right REDUCE
|
||||
if mval.arg >= 0 or loop_start.arg != 0:
|
||||
# TODO: support and test this with other mvals and loop_starts
|
||||
if DEBUG >= 1: print(f"WARNING, NOT FOLDING: mval:{mval.arg} loop_start:{loop_start.arg}")
|
||||
return None
|
||||
comprange = UOp.min(loop_end, UOp.max(UOp.alu(BinaryOps.IDIV, idx-compval-mval, mval) + (loop_end-loop_start), loop_start))
|
||||
return UOp(UOps.UNMUL, multconst.dtype, (comprange.cast(multconst.dtype) * multconst, loop_end-loop_start))
|
||||
return UOp(UOps.REDUCE, reduce_allow_any_len.dtype, (comprange.cast(multconst.dtype) * multconst,) +
|
||||
tuple(x for x in reduce_allow_any_len.src[1:] if x is not rng), reduce_allow_any_len.arg)
|
||||
|
||||
def index_collapse(idx,rng,buf,add,mul,ld,reduce_allow_any_len):
|
||||
if rng not in reduce_allow_any_len.src: return None
|
||||
return UOp(reduce_allow_any_len.op, reduce_allow_any_len.dtype, (UOp(ld.op, ld.dtype, (buf, add+mul*idx)),)+
|
||||
tuple(x for x in reduce_allow_any_len.src[1:] if x is not rng), reduce_allow_any_len.arg)
|
||||
|
||||
# this is symbolic 2.0
|
||||
constant_folder = PatternMatcher([
|
||||
@@ -154,29 +160,24 @@ constant_folder = PatternMatcher([
|
||||
lambda add, wmma: UOp(wmma.op, wmma.dtype, (wmma.src[0], wmma.src[1], wmma.src[2]+add), wmma.arg)),
|
||||
# threefry
|
||||
(UOp(UOps.ALU, dtype=dtypes.uint64, src=(UOp.var("x"), UOp.var("seed")), arg=BinaryOps.THREEFRY), threefry2x32),
|
||||
# arange loop folding (early)
|
||||
((UOp.var("idx") + UOp.cvar("mval") * UOp(UOps.RANGE, src=(UOp.var("loop_start"), UOp.var("loop_end"))).name("rng")).lt(UOp.cvar("compval")).where(
|
||||
UOp.cvar("multconst"), UOp.const(None, 0)), loop_collapse),
|
||||
((UOp.var("idx") - UOp(UOps.RANGE, src=(UOp.var("loop_start"), UOp.var("loop_end"))).name("rng")).lt(UOp.cvar("compval")).where(
|
||||
UOp.cvar("multconst"), UOp.const(None, 0)), lambda **kwargs: loop_collapse(mval=UOp.const(dtypes.int, -1), **kwargs)),
|
||||
# sum collapse to mul (with possible GEP)
|
||||
(UPat(UOps.PHI, src=(UPat(UOps.DEFINE_ACC, name="phi_input", src=[UPat(UOps.CONST), UPat(UOps.RANGE, name="loop")]),
|
||||
UPat(UOps.ALU, BinaryOps.ADD, src=(UPat(name="val1"), UPat(name="val2"))))), sum_collapse),
|
||||
(UPat(UOps.PHI, src=(UPat(UOps.GEP, name="phi_input", src=(UPat(UOps.DEFINE_ACC, src=[UPat(UOps.CONST), UPat(UOps.RANGE, name="loop")]),)),
|
||||
UPat(UOps.ALU, BinaryOps.ADD, src=(UPat(name="val1"), UPat(name="val2"))))), sum_collapse),
|
||||
# deal with UNMUL
|
||||
(UOp.cvar('c1') * UOp(UOps.UNMUL, src=(UOp.cvar('c2'), UOp.var('v'))), lambda c1,c2,v: v if c1.arg == c2.arg else None),
|
||||
(UOp.cvar('c1') * (UOp.var('add') + UOp(UOps.UNMUL, src=(UOp.cvar('c2'), UOp.var('v')))),
|
||||
lambda c1, add, c2, v: (add*c1+v) if c1.arg == c2.arg else None),
|
||||
(UOp(UOps.UNMUL, src=(UOp.const(None, 0).name('zero'), UOp.var())), lambda zero: zero),
|
||||
(UOp(UOps.UNMUL).name('unmul').cast().name('root'), lambda root,unmul: UOp(UOps.UNMUL, root.dtype, (unmul.src[0].cast(root.dtype), unmul.src[1]))),
|
||||
# arange loop folding (reduce)
|
||||
(UOp(UOps.REDUCE, src=((UOp.var("idx") + UOp.cvar("mval") * UOp(UOps.RANGE, src=(UOp.var("loop_start"), UOp.var("loop_end"))).name("rng"))
|
||||
.lt(UOp.cvar("compval")).where(UOp.cvar("multconst"), UOp.const(None, 0)),), arg=ReduceOps.SUM).name("reduce_allow_any_len"), loop_collapse),
|
||||
(UOp(UOps.REDUCE, src=((UOp.var("idx") - UOp(UOps.RANGE, src=(UOp.var("loop_start"), UOp.var("loop_end"))).name("rng"))
|
||||
.lt(UOp.cvar("compval")).where(UOp.cvar("multconst"), UOp.const(None, 0)),), arg=ReduceOps.SUM).name("reduce_allow_any_len"),
|
||||
lambda **kwargs: loop_collapse(mval=UOp.const(dtypes.int, -1), **kwargs)),
|
||||
# indexing (with a multiply offset)!
|
||||
(UOp.var('idx').eq(UOp(UOps.RANGE).name("rng")).cast()*
|
||||
UOp(UOps.LOAD, src=(UOp.var("buf"), UOp.var('add')+UOp.var('mul')*UOp(UOps.RANGE).name("rng"))).name("ld"),
|
||||
lambda idx,rng,buf,add,mul,ld: UOp(UOps.UNMUL, ld.dtype, (UOp(ld.op, ld.dtype, (buf, add+mul*idx)), rng.src[1]-rng.src[0]))),
|
||||
(UOp.var('idx').eq(UOp(UOps.RANGE).name("rng")).where(
|
||||
UOp(UOps.LOAD, src=(UOp.var("buf"), UOp.var('add')+UOp.var('mul')*UOp(UOps.RANGE).name("rng"))).name("ld"), UOp.const(None, 0.0)),
|
||||
lambda idx,rng,buf,add,mul,ld: UOp(UOps.UNMUL, ld.dtype, (UOp(ld.op, ld.dtype, (buf, add+mul*idx)), rng.src[1]-rng.src[0]))),
|
||||
(UOp(UOps.REDUCE, src=(UOp.var('idx').eq(UOp(UOps.RANGE).name("rng")).cast()*
|
||||
UOp(UOps.LOAD, src=(UOp.var("buf"), UOp.var('add')+UOp.var('mul')*UOp(UOps.RANGE).name("rng"))).name("ld"),),
|
||||
arg=ReduceOps.SUM).name("reduce_allow_any_len"), index_collapse),
|
||||
(UOp(UOps.REDUCE, src=(UOp.var('idx').eq(UOp(UOps.RANGE).name("rng")).where(
|
||||
UOp(UOps.LOAD, src=(UOp.var("buf"), UOp.var('add')+UOp.var('mul')*UOp(UOps.RANGE).name("rng"))).name("ld"), UOp.const(None, 0.0)),),
|
||||
arg=ReduceOps.SUM).name("reduce_allow_any_len"), index_collapse),
|
||||
# other arange folders
|
||||
(UOp.cvar("c1") - (UOp.var("x") + UOp.cvar("c2")), lambda c1, c2, x: (c1-c2)-x), # c1 - (x + c2) -> (c1-c2) - x
|
||||
# max on special can go away (TODO: special should be variable, same thing applies)
|
||||
@@ -539,7 +540,7 @@ class UOpGraph:
|
||||
for u, x in scope_end.items(): self._uops.insert(self._uops.index(x)+1, UOp(END_FOR_UOP[u.op][1], None, (u,)))
|
||||
|
||||
# sanity checks (NOTE: these can cause things to be skipped in BEAM)
|
||||
bad_ops = dedup([x.op for x in self._uops if x.op in {UOps.EXPAND, UOps.CONTRACT, UOps.REDUCE, UOps.UNMUL}])
|
||||
bad_ops = dedup([x.op for x in self._uops if x.op in {UOps.EXPAND, UOps.CONTRACT, UOps.REDUCE}])
|
||||
try:
|
||||
type_verify(self.uops)
|
||||
assert self._uops[-1].op is UOps.SINK, f"didn't end with SINK, ended with {self._uops[-1]}"
|
||||
|
||||
@@ -15,7 +15,7 @@ class UOps(Enum):
|
||||
SINK = auto(); VAR = auto(); EXPAND = auto(); CONTRACT = auto() # noqa: E702
|
||||
DEFINE_GLOBAL = auto(); DEFINE_VAR = auto(); DEFINE_LOCAL = auto(); DEFINE_ACC = auto() # noqa: E702
|
||||
CONST = auto(); SPECIAL = auto() # noqa: E702
|
||||
NOOP = auto(); UNMUL = auto(); GEP = auto() # noqa: E702
|
||||
NOOP = auto(); GEP = auto() # noqa: E702
|
||||
# math ops
|
||||
CAST = auto(); BITCAST = auto(); VECTORIZE = auto() # noqa: E702
|
||||
ALU = auto(); REDUCE = auto(); WMMA = auto() # noqa: E702
|
||||
@@ -35,7 +35,7 @@ class UOp:
|
||||
src: Tuple[UOp, ...] = tuple()
|
||||
arg: Any = None
|
||||
def commutative(self) -> bool:
|
||||
return self.op is UOps.UNMUL or (self.op is UOps.ALU and \
|
||||
return (self.op is UOps.ALU and \
|
||||
self.arg in {BinaryOps.ADD, BinaryOps.MUL, BinaryOps.MAX, BinaryOps.CMPNE, BinaryOps.XOR, BinaryOps.AND, BinaryOps.OR})
|
||||
@functools.cached_property
|
||||
def cmp_tuple(self):
|
||||
|
||||
Reference in New Issue
Block a user