diff --git a/test/test_jit.py b/test/test_jit.py index 7c6bf77206..a9a29ed90d 100644 --- a/test/test_jit.py +++ b/test/test_jit.py @@ -277,6 +277,38 @@ class TestJit(unittest.TestCase): assert len(res3) == 5, "All values should be different, rand works in jit." assert res3 != res2, "Jit rand is diff with diff seeds" + @unittest.expectedFailure # TODO: fix + def test_jit_v_nojit_random_regen(self): + def f(a, b): + rn = Tensor.randn(*a.shape) + rn = rn * a + rn2 = Tensor.randn(*a.shape) + rn2 = rn2 * b + rn = rn + rn2 + rn2 = rn2 + Tensor.randn(*a.shape) + return ((a+b)*rn).realize(), ((a+b)*rn2).realize() + Tensor.manual_seed(0) + a = Tensor.randn(10, 10).realize() # realize these before resetting the random seed + b = Tensor.randn(10, 10).realize() + + Tensor.manual_seed(1234) + without_jit = set() + for _ in range(5): + o1, o2 = f(a, b) + without_jit.add(o1.numpy()[0][0]) + without_jit.add(o2.numpy()[0][0]) + assert len(without_jit) == 10, "All values should be different." + + Tensor.manual_seed(1234) + jf = TinyJit(f) + with_jit = set() + for _ in range(5): + o1, o2 = jf(a, b) + with_jit.add(o1.numpy()[0][0]) + with_jit.add(o2.numpy()[0][0]) + assert len(with_jit) == 10, "All values should be different." + assert with_jit == without_jit, "Jit rand produced different values from no jit." + def test_jit_multiple_random_regen(self): def f(a, b): rn = Tensor.randn(*a.shape) diff --git a/tinygrad/codegen/symbolic.py b/tinygrad/codegen/symbolic.py index e25477388c..ab4c83b914 100644 --- a/tinygrad/codegen/symbolic.py +++ b/tinygrad/codegen/symbolic.py @@ -314,11 +314,11 @@ def simplify_valid(valid:UOp) -> UOp|None: # ***** threefry ***** def threefry2x32(x: UOp, key: UOp): - # split x into two uint32, since x in a uint64 + # split x and key from uint64 to two uint32 x0, x1 = (x & 0xffffffff).cast(dtypes.uint32), ((x // 2**32) & 0xffffffff).cast(dtypes.uint32) + key0, key1 = (key & 0xffffffff).cast(dtypes.uint32), ((key // 2**32) & 0xffffffff).cast(dtypes.uint32) rotations = [[13, 15, 26, 6], [17, 29, 16, 24]] - key0, key1 = (key & 0xffffffff).cast(dtypes.uint32), ((key // 2**32) & 0xffffffff).cast(dtypes.uint32) ks = [key1, key0 ^ key1 ^ 0x1BD11BDA, key0] xr = [x0 + ks[-1], x1 + ks[0]] for i in range(5):