From 99f6d89dfbcd7ef947b30958711221075cd6dc7e Mon Sep 17 00:00:00 2001 From: chenyu Date: Thu, 8 May 2025 22:38:56 -0400 Subject: [PATCH] tighter idiv bound for symbolic denominator (#10226) --- test/unit/test_uop_vmin_vmax.py | 21 +++++++++++++++++++++ tinygrad/ops.py | 7 ++++--- 2 files changed, 25 insertions(+), 3 deletions(-) diff --git a/test/unit/test_uop_vmin_vmax.py b/test/unit/test_uop_vmin_vmax.py index 59ce0e9a93..e0467e3b7f 100644 --- a/test/unit/test_uop_vmin_vmax.py +++ b/test/unit/test_uop_vmin_vmax.py @@ -147,6 +147,27 @@ class TestVminVmaxDivMod(unittest.TestCase): self.assertEqual(uop.vmin, -3) self.assertEqual(uop.vmax, 3) + def test_vmin_vmax_div_symbolic(self): + x = UOp.variable('x', 1, 10) + y = UOp.variable('y', 3, 5) + self.assertEqual((x//y).vmin, 0) + self.assertEqual((x//y).vmax, 3) + self.assertEqual(((-x)//y).vmin, -3) + self.assertEqual(((-x)//y).vmax, 0) + self.assertEqual((x//(-y)).vmin, -3) + self.assertEqual((x//(-y)).vmax, 0) + self.assertEqual(((-x)//(-y)).vmin, 0) + self.assertEqual(((-x)//(-y)).vmax, 3) + + self.assertEqual((100//y).vmin, 20) + self.assertEqual((100//y).vmax, 33) + self.assertEqual(((-100)//y).vmin, -33) + self.assertEqual(((-100)//y).vmax, -20) + self.assertEqual((100//(-y)).vmin, -33) + self.assertEqual((100//(-y)).vmax, -20) + self.assertEqual(((-100)//(-y)).vmin, 20) + self.assertEqual(((-100)//(-y)).vmax, 33) + def test_vmin_vmax_mod_positive(self): # vmin and vmax for modulo of a variable by a positive constant x = UOp.variable('x', 10, 20) diff --git a/tinygrad/ops.py b/tinygrad/ops.py index 2ee80b53b2..d693567f7e 100644 --- a/tinygrad/ops.py +++ b/tinygrad/ops.py @@ -634,9 +634,10 @@ class UOp(MathTrait, metaclass=UOpMetaClass): if (c:=s1_vmin) == s1_vmax: # s1 is a const if c > 0: return cdiv(s0_vmin, c), cdiv(s0_vmax, c) if c < 0: return cdiv(s0_vmax, c), cdiv(s0_vmin, c) - # don't know exact bounds, but know the sign - if (s0_vmax <= 0 and s1_vmax < 0) or (s0_vmin >= 0 and s1_vmin > 0): return 0, dtypes.max(self.dtype) - if (s0_vmax <= 0 and s1_vmin > 0) or (s0_vmin >= 0 and s1_vmax < 0): return dtypes.min(self.dtype), 0 + if (s0_vmax <= 0 and s1_vmax < 0): return cdiv(s0_vmax, s1_vmin), cdiv(s0_vmin, s1_vmax) + if (s0_vmin >= 0 and s1_vmin > 0): return cdiv(s0_vmin, s1_vmax), cdiv(s0_vmax, s1_vmin) + if (s0_vmax <= 0 and s1_vmin > 0): return cdiv(s0_vmin, s1_vmin), cdiv(s0_vmax, s1_vmax) + if (s0_vmin >= 0 and s1_vmax < 0): return cdiv(s0_vmax, s1_vmax), cdiv(s0_vmin, s1_vmin) if self.op is Ops.MAX: return max(s0_vmin, s1_vmin), max(s0_vmax, s1_vmax) if self.op is Ops.CMPLT: return (s0_vmax