forked from tinygrad/tinygrad
Merge remote-tracking branch 'origin/master' into qwen36_27b_amd_900
# Conflicts: # tinygrad/llm/model.py
This commit is contained in:
@@ -307,8 +307,8 @@ class TestRecurse(unittest.TestCase):
|
||||
def test_inf_loop(self):
|
||||
a = UOp.const(3)
|
||||
pm = PatternMatcher([
|
||||
(UPat(Ops.CONST, arg=3, name="x"), lambda x: x.replace(arg=4)),
|
||||
(UPat(Ops.CONST, arg=4, name="x"), lambda x: x.replace(arg=3)),
|
||||
(UPat(Ops.CONST, arg=3, name="x"), lambda x: UOp.const(4, x.dtype)),
|
||||
(UPat(Ops.CONST, arg=4, name="x"), lambda x: UOp.const(3, x.dtype)),
|
||||
])
|
||||
with self.assertRaises(RuntimeError):
|
||||
graph_rewrite(a, pm)
|
||||
@@ -316,8 +316,8 @@ class TestRecurse(unittest.TestCase):
|
||||
def test_inf_loop_bottom_up(self):
|
||||
a = UOp.const(3)
|
||||
pm = PatternMatcher([
|
||||
(UPat(Ops.CONST, arg=3, name="x"), lambda x: x.replace(arg=4)),
|
||||
(UPat(Ops.CONST, arg=4, name="x"), lambda x: x.replace(arg=3)),
|
||||
(UPat(Ops.CONST, arg=3, name="x"), lambda x: UOp.const(4, x.dtype)),
|
||||
(UPat(Ops.CONST, arg=4, name="x"), lambda x: UOp.const(3, x.dtype)),
|
||||
])
|
||||
with self.assertRaises(RuntimeError):
|
||||
graph_rewrite(a, pm, bottom_up=True)
|
||||
@@ -378,8 +378,8 @@ class TestWalkRewrite(unittest.TestCase):
|
||||
"""A bouncing pattern applies once and stops instead of looping."""
|
||||
a = UOp.const(3)
|
||||
pm = PatternMatcher([
|
||||
(UPat(Ops.CONST, arg=3, name="x"), lambda x: x.replace(arg=4)),
|
||||
(UPat(Ops.CONST, arg=4, name="x"), lambda x: x.replace(arg=3)),
|
||||
(UPat(Ops.CONST, arg=3, name="x"), lambda x: UOp.const(4, x.dtype)),
|
||||
(UPat(Ops.CONST, arg=4, name="x"), lambda x: UOp.const(3, x.dtype)),
|
||||
])
|
||||
with self.assertRaises(RuntimeError):
|
||||
graph_rewrite(a, pm, bottom_up=True)
|
||||
@@ -456,8 +456,8 @@ class TestWalkRewrite(unittest.TestCase):
|
||||
"""Bottom-up walk also applies once per node, no fixed-point iteration."""
|
||||
a = UOp.const(3)
|
||||
pm = PatternMatcher([
|
||||
(UPat(Ops.CONST, arg=3, name="x"), lambda x: x.replace(arg=4)),
|
||||
(UPat(Ops.CONST, arg=4, name="x"), lambda x: x.replace(arg=3)),
|
||||
(UPat(Ops.CONST, arg=3, name="x"), lambda x: UOp.const(4, x.dtype)),
|
||||
(UPat(Ops.CONST, arg=4, name="x"), lambda x: UOp.const(3, x.dtype)),
|
||||
])
|
||||
ret = graph_rewrite(a, pm, bottom_up=True, walk=True)
|
||||
self.assertIs(ret, UOp.const(4))
|
||||
@@ -511,7 +511,7 @@ class TestWalkRewrite(unittest.TestCase):
|
||||
def bpm_match(ctx, x):
|
||||
ctx.append((x.val if x.op is Ops.CONST else x.op, "bpm"))
|
||||
# rewrite const(1) -> const(10), short-circuiting its subtree
|
||||
if x.op is Ops.CONST and x.val == 1: return x.replace(arg=10)
|
||||
if x.op is Ops.CONST and x.val == 1: return UOp.const(10, x.dtype)
|
||||
return None
|
||||
def pm_match(ctx, x):
|
||||
ctx.append((x.val if x.op is Ops.CONST else x.op, "pm"))
|
||||
|
||||
@@ -593,7 +593,7 @@ class TestUOpTags(unittest.TestCase):
|
||||
def test_inc_by_one(self):
|
||||
g = UOp.const(1) + UOp.const(1)
|
||||
assert g.ssimplify() == 2
|
||||
pm_plus_1 = PatternMatcher([(UPat(Ops.CONST, name="x"), lambda x: x.replace(arg=x.val+1, tag=1) if x.tag is None else None)])
|
||||
pm_plus_1 = PatternMatcher([(UPat(Ops.CONST, name="x"), lambda x: UOp.const(x.val+1, x.dtype).rtag(1) if x.tag is None else None)])
|
||||
pm_strip_tags = PatternMatcher([(UPat(GroupOp.All, name="x"), lambda x: x.replace(tag=None) if x.tag is not None else None)])
|
||||
g = graph_rewrite(g, pm_plus_1)
|
||||
assert g.ssimplify() == 4
|
||||
|
||||
@@ -126,6 +126,12 @@ class TestConstFloatEq(unittest.TestCase):
|
||||
self.assertFalse(nan == Invalid)
|
||||
self.assertTrue(nan != Invalid) # __ne__ must defer to the reflected eq, not swallow NotImplemented
|
||||
|
||||
def test_invalid_eq_defers_to_reflected(self):
|
||||
class HoldsInvalid: # a carrier that knows it holds Invalid. returning False for foreign types would silence its eq
|
||||
def __eq__(self, other): return other is Invalid
|
||||
self.assertTrue(Invalid == HoldsInvalid())
|
||||
self.assertFalse(Invalid != HoldsInvalid())
|
||||
|
||||
def test_matchers_agree_on_nan(self):
|
||||
n = UOp.const(math.nan, dtypes.float32)
|
||||
for compiled in (False, True):
|
||||
@@ -447,7 +453,7 @@ class TestUPatHelpers(unittest.TestCase):
|
||||
|
||||
class TestUopsObject(unittest.TestCase):
|
||||
def test_timing(self):
|
||||
with Timing("create 10k uops:"): ret = [UOp(Ops.CONST, dtypes.int, arg=10000000+i) for i in range(10000)]
|
||||
with Timing("create 10k uops:"): ret = [UOp.const(10000000+i, dtypes.int) for i in range(10000)]
|
||||
assert len(ret) == 10000
|
||||
|
||||
def test_nested(self):
|
||||
|
||||
@@ -147,21 +147,21 @@ class TestUOpsStats(unittest.TestCase):
|
||||
#MULACC should have the same stats as MUL + ADD
|
||||
def test_mulacc(self):
|
||||
globl = UOp.param(0, dtypes.int, (3,))
|
||||
o1 = UOp(Ops.CONST, dtypes.int, tuple(), 1)
|
||||
o2 = UOp(Ops.CONST, dtypes.int, tuple(), 2)
|
||||
o1 = UOp.const(1, dtypes.int)
|
||||
o2 = UOp.const(2, dtypes.int)
|
||||
u1 = globl.index(o1)
|
||||
u2 = globl.index(o2)
|
||||
u3 = UOp(Ops.CONST, dtypes.int, tuple(), 3)
|
||||
u3 = UOp.const(3, dtypes.int)
|
||||
u4 = UOp(Ops.MUL, src=(u1,u2))
|
||||
u5 = UOp(Ops.ADD, src=(u4,u3))
|
||||
uops = tuple(u5.toposort())
|
||||
|
||||
globl = UOp.param(0, dtypes.int, (3,))
|
||||
o1 = UOp(Ops.CONST, dtypes.int, tuple(), 1)
|
||||
o2 = UOp(Ops.CONST, dtypes.int, tuple(), 2)
|
||||
o1 = UOp.const(1, dtypes.int)
|
||||
o2 = UOp.const(2, dtypes.int)
|
||||
u1 = globl.index(o1)
|
||||
u2 = globl.index(o2)
|
||||
u3 = UOp(Ops.CONST, dtypes.int, tuple(), 3)
|
||||
u3 = UOp.const(3, dtypes.int)
|
||||
u4 = UOp(Ops.MULACC, src=(u1,u2,u3))
|
||||
uops_fma = tuple(u4.toposort())
|
||||
|
||||
|
||||
@@ -97,7 +97,7 @@ class TestViz(unittest.TestCase):
|
||||
# VIZ tracks rewrites up to and including the error
|
||||
def count_3(x:UOp):
|
||||
assert x.val <= 3
|
||||
return x.replace(arg=x.val+1)
|
||||
return UOp.const(x.val+1, x.dtype)
|
||||
err_pm = PatternMatcher([(UPat.cvar("x"), count_3),])
|
||||
a = UOp.const(1)
|
||||
with save_viz() as viz:
|
||||
@@ -202,8 +202,8 @@ class TestViz(unittest.TestCase):
|
||||
a = UOp.const(3)
|
||||
b = UOp.const(4)
|
||||
pm = PatternMatcher([
|
||||
(UPat(Ops.CONST, arg=3, name="x"), lambda x: x.replace(arg=4)),
|
||||
(UPat(Ops.CONST, arg=4, name="x"), lambda x: x.replace(arg=3)),
|
||||
(UPat(Ops.CONST, arg=3, name="x"), lambda x: UOp.const(4, x.dtype)),
|
||||
(UPat(Ops.CONST, arg=4, name="x"), lambda x: UOp.const(3, x.dtype)),
|
||||
])
|
||||
with save_viz() as viz:
|
||||
# use smaller stack limit for faster test (default is 250000)
|
||||
@@ -224,7 +224,7 @@ class TestViz(unittest.TestCase):
|
||||
list(viz.get_details(0, 0))
|
||||
|
||||
def test_enter_calls_rewrite(self):
|
||||
pm = PatternMatcher([(UPat(Ops.CONST, arg=3, name="x"), lambda x: x.replace(arg=4))])
|
||||
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,)),))
|
||||
|
||||
Reference in New Issue
Block a user