From 407ca543823f91ae720a940c4e96b502f2a1bdcd Mon Sep 17 00:00:00 2001 From: chenyu Date: Sat, 5 Apr 2025 05:12:17 -0400 Subject: [PATCH] symbolic fold double where (#9436) * symbolic fold double where a.where(b.where(c, d), d) -> (a & b).where(c, d). a pattern in optimizer * test case --- test/unit/test_uop_symbolic.py | 17 +++++++++++++++++ tinygrad/codegen/symbolic.py | 2 ++ 2 files changed, 19 insertions(+) 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 **