From 2b4375f36d7e01f0f8a15c382c87b582991d8bf7 Mon Sep 17 00:00:00 2001 From: Sieds Lykles <93992551+S-Lykles@users.noreply.github.com> Date: Wed, 21 May 2025 12:46:13 +0200 Subject: [PATCH] Correct divmod folding behind flag (#10433) * add flag * add test * remove import --- test/unit/test_uop_symbolic.py | 9 +++++++++ tinygrad/codegen/symbolic.py | 15 ++++++++------- tinygrad/helpers.py | 1 + 3 files changed, 18 insertions(+), 7 deletions(-) diff --git a/test/unit/test_uop_symbolic.py b/test/unit/test_uop_symbolic.py index f8482f2468..a70e070c15 100644 --- a/test/unit/test_uop_symbolic.py +++ b/test/unit/test_uop_symbolic.py @@ -4,6 +4,7 @@ import unittest, pickle, functools from tinygrad.dtype import dtypes, ConstType from tinygrad.codegen import full_rewrite from tinygrad.codegen.devectorizer import sym +from tinygrad.helpers import Context from tinygrad.uop.ops import UOp, Ops, graph_rewrite, sym_infer from tinygrad import Variable @@ -322,9 +323,17 @@ class TestSymbolic(unittest.TestCase): def test_div_cancel(self): self.helper_test_variable(usum([uconst(-40), Variable("a", 0, 10)*2, Variable("b", 0, 10)*40])//40, -1, 9, "(b+-1)") + def test_div_cancel_correct(self): + with Context(CORRECT_DIVMOD_FOLDING=1): + self.helper_test_variable(usum([uconst(-40), Variable("a", 0, 10)*2, Variable("b", 0, 10)*40])//40, -1, 9, "(((a+(b*20))+-20)//20)") + def test_mod_cancel(self): self.helper_test_variable(usum([uconst(-40), Variable("a", 0, 10)*2, Variable("b", 0, 10)*40]) % 40, 0, 20, "(a*2)") + def test_mod_cancel_correct(self): + with Context(CORRECT_DIVMOD_FOLDING=1): + self.helper_test_variable(usum([uconst(-40), Variable("a", 0, 10)*2, Variable("b", 0, 10)*40]) % 40, -38, 38, "((((a+(b*20))+-20)%20)*2)") + def test_mul_div(self): self.helper_test_variable((Variable("a", 0, 10)*4)//4, 0, 10, "a") diff --git a/tinygrad/codegen/symbolic.py b/tinygrad/codegen/symbolic.py index d240aa3668..1778e2316b 100644 --- a/tinygrad/codegen/symbolic.py +++ b/tinygrad/codegen/symbolic.py @@ -4,7 +4,7 @@ import math, operator, struct, functools from collections import defaultdict from tinygrad.uop.ops import Ops, PatternMatcher, UPat, UOp, GroupOp, exec_alu from tinygrad.dtype import ConstType, dtypes, PtrDType -from tinygrad.helpers import partition, all_same, prod, flatten, get_single_element, cdiv, cmod +from tinygrad.helpers import partition, all_same, prod, flatten, get_single_element, cdiv, cmod, CORRECT_DIVMOD_FOLDING from tinygrad.codegen.transcendental import xpow # ******** phase 1 of symbolic used to live in ops, it's the most generic folding rules ******** @@ -152,12 +152,13 @@ def div_and_mod_folding(x: UOp, y: UOp, which: Literal[Ops.MOD, Ops.IDIV], split y2 = cmod(factors[0]*v.vmax+const, c) if which is Ops.MOD else cdiv(factors[0]*v.vmax+const, c) return (y2-y1)*(v-v.vmin) + y1 - # a//c = (a-a%c)/c, if we can fold a%c, we can fold a//c - # within a mod we can freely subtract multiples of c, we use this to see if a is congruent to an expression whose vmin/vmax are between 0 and c - rems = [min(r, r-c, key=abs) for r in remainders] - if (rem:=sum(r*v for r,v in zip(rems,svars))+const%c).vmin//c==rem.vmax//c and all(f > 0 for f in factors): - if which is Ops.MOD: return rem - rem.vmin//c*c - return sum((f-r)//c * v for f,r,v in zip(factors,rems,svars)) + (const-const%c+rem.vmin//c*c)//c + if not CORRECT_DIVMOD_FOLDING or x_min>=0: + # a//c = (a-a%c)/c, if we can fold a%c, we can fold a//c + # within a mod we can freely subtract multiples of c, we use this to see if a is congruent to an expression whose vmin/vmax are between 0 and c + rems = [min(r, r-c, key=abs) for r in remainders] + if (rem:=sum(r*v for r,v in zip(rems,svars))+const%c).vmin//c==rem.vmax//c and all(f > 0 for f in factors): + if which is Ops.MOD: return rem - rem.vmin//c*c + return sum((f-r)//c * v for f,r,v in zip(factors,rems,svars)) + (const-const%c+rem.vmin//c*c)//c if (g:=math.gcd(gcd, const))!=1: ret = UOp(which, x.dtype, src=(sum(f//g * v for f,v in zip(factors, svars)) + const//g, x.const_like(c//g))) diff --git a/tinygrad/helpers.py b/tinygrad/helpers.py index 41217d3e71..da28675d09 100644 --- a/tinygrad/helpers.py +++ b/tinygrad/helpers.py @@ -118,6 +118,7 @@ CACHELEVEL, IGNORE_BEAM_CACHE, DEVECTORIZE = ContextVar("CACHELEVEL", 2), Contex DISABLE_COMPILER_CACHE = ContextVar("DISABLE_COMPILER_CACHE", 0) DONT_REALIZE_EXPAND, DONT_GROUP_REDUCES = ContextVar("DONT_REALIZE_EXPAND", 0), ContextVar("DONT_GROUP_REDUCES", 0) QUANTIZE, VALIDATE_WITH_CPU, IGNORE_OOB = ContextVar("QUANTIZE", 0), ContextVar("VALIDATE_WITH_CPU", 0), ContextVar("IGNORE_OOB", 1) +CORRECT_DIVMOD_FOLDING = ContextVar("CORRECT_DIVMOD_FOLDING", 0) @dataclass(frozen=True) class Metadata: