diff --git a/test/unit/test_uop_symbolic.py b/test/unit/test_uop_symbolic.py index cf8ba3237a..c76cd7a91e 100644 --- a/test/unit/test_uop_symbolic.py +++ b/test/unit/test_uop_symbolic.py @@ -555,6 +555,23 @@ class TestSymbolic(unittest.TestCase): self.assertEqual(rewritten_uop, cond.where(a.cast(dtypes.half), b.cast(dtypes.half))) + def test_where_merge_branches(self): + cond1 = Variable("s", 0, 10) < 6 + cond2 = Variable("s", 0, 10) > 2 + a = Variable("a", 0, 3) + b = Variable("b", 0, 3) + expr = cond1.where(cond2.where(a, b), b) + self.helper_test_variable(expr, 0, 3, "(a if ((s<6)&(2 (a if (s<5) else b) + self.helper_test_variable(expr, 0, 3, "(a if (s<5) else b)") + def test_symbolic_div(self): # from symbolic arange a = Variable("a", 1, 10) diff --git a/tinygrad/codegen/symbolic.py b/tinygrad/codegen/symbolic.py index 91758b3b9d..0b2db0178c 100644 --- a/tinygrad/codegen/symbolic.py +++ b/tinygrad/codegen/symbolic.py @@ -461,6 +461,8 @@ sym = symbolic_flat+PatternMatcher([ # ** where ** # push cast to branches (UPat.var("s").where(UPat.var("a"), UPat.var("b")).cast().named("cast"), lambda s,a,b,cast: s.where(a.cast(cast.dtype), b.cast(cast.dtype))), + # a.where(b.where(c, d), d) -> (a & b).where(c, d) + (UPat.var("a").where(UPat.var("b").where(UPat.var("c"), UPat.var("d")), UPat.var("d")), lambda a,b,c,d: (a&b).where(c,d)), # ** pow ** ((UPat(Ops.POW, name="p"), lambda p: xpow(*p.src))), # ** load/store folding **