From 98cd9e8926fb03da962b33437493f6ddfe4d6f17 Mon Sep 17 00:00:00 2001 From: Paul Gustafson Date: Mon, 27 Nov 2023 18:37:44 -0800 Subject: [PATCH] Add assertion to prevent nonsense mod values (#2474) --- test/unit/test_symbolic.py | 6 ++++++ tinygrad/shape/symbolic.py | 2 +- 2 files changed, 7 insertions(+), 1 deletion(-) diff --git a/test/unit/test_symbolic.py b/test/unit/test_symbolic.py index 673657d1d1..03750d8acf 100644 --- a/test/unit/test_symbolic.py +++ b/test/unit/test_symbolic.py @@ -402,6 +402,12 @@ class TestSymbolicSymbolicOps(unittest.TestCase): with self.assertRaises(AssertionError): lt3 = (x < 3) + def test_nested_variable_mod(self): + i = Variable("i", 1, 5) + idx0 = Variable("idx0", 0, i) + with self.assertRaises(AssertionError): + assert idx0 % 2 == idx0 + def test_num_node_mul_node(self): a = Variable("a", 1, 5) b = NumNode(2) * a diff --git a/tinygrad/shape/symbolic.py b/tinygrad/shape/symbolic.py index 87062e012f..f8813d09a5 100644 --- a/tinygrad/shape/symbolic.py +++ b/tinygrad/shape/symbolic.py @@ -94,7 +94,7 @@ class Node: if self == b: return NumNode(0) if (b - self).min > 0 and self.min >= 0: return self # b - self simplifies the node raise RuntimeError(f"not supported: {self} % {b}") - assert b > 0 + assert b > 0 and isinstance(self.max, int) and isinstance(self.min, int), 'not supported' if b == 1: return NumNode(0) if self.min >= 0 and self.max < b: return self if (self.min//b) == (self.max//b): return self - (b*(self.min//b))