diff --git a/tinygrad/codegen/decomp/op.py b/tinygrad/codegen/decomp/op.py index 6b0526b498..6a48cdca53 100644 --- a/tinygrad/codegen/decomp/op.py +++ b/tinygrad/codegen/decomp/op.py @@ -47,17 +47,17 @@ def fast_idiv(ren: Renderer, x: UOp, d: int, dont_cast=False) -> UOp|None: def threefry2x32(x: UOp, key: UOp): # 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) + x0, x1 = x.cast(dtypes.uint32), (x >> 32).cast(dtypes.uint32) + key0, key1 = key.cast(dtypes.uint32), (key >> 32).cast(dtypes.uint32) rotations = [[13, 15, 26, 6], [17, 29, 16, 24]] ks = [key1, key0 ^ key1 ^ 0x1BD11BDA, key0] xr:list[UOp] = [x0 + ks[-1], x1 + ks[0]] for i in range(5): - for r in rotations[i % 2]: xr[0], xr[1] = (x0 := xr[0] + xr[1]), x0 ^ ((xr[1] * 2**r) + (xr[1] // 2**(32 - r))) + for r in rotations[i % 2]: xr[0], xr[1] = (x0 := xr[0] + xr[1]), x0 ^ ((xr[1] << r) + (xr[1] >> (32 - r))) xr = [(xr[0] + ks[i % 3]), (xr[1] + ks[(i + 1) % 3] + i + 1)] - return xr[1].cast(dtypes.uint64) * 2**32 | xr[0].cast(dtypes.uint64) + return (xr[1].cast(dtypes.uint64) << 32) | xr[0].cast(dtypes.uint64) # ***** decomposition patterns ***** diff --git a/tinygrad/mixin/rand.py b/tinygrad/mixin/rand.py index 122f4e1f41..e1e09fd94d 100644 --- a/tinygrad/mixin/rand.py +++ b/tinygrad/mixin/rand.py @@ -12,7 +12,7 @@ class RandMixin(OpMixin): def _threefry_random_bits(key, counts0, counts1): x = (counts1.cast(dtypes.uint64) << 32) | counts0.cast(dtypes.uint64) x = x.threefry((key[1].cast(dtypes.uint64) << 32) | key[0].cast(dtypes.uint64)) - return (x & 0xffffffff).cast(dtypes.uint32).cat(((x >> 32) & 0xffffffff).cast(dtypes.uint32)) + return x.cast(dtypes.uint32).cat((x >> 32).cast(dtypes.uint32)) @classmethod def random_bits(cls, key:Self, counter:Self, num:int) -> Self: diff --git a/tinygrad/uop/symbolic.py b/tinygrad/uop/symbolic.py index 8830abcd70..76c600f90d 100644 --- a/tinygrad/uop/symbolic.py +++ b/tinygrad/uop/symbolic.py @@ -127,7 +127,6 @@ symbolic_simple = pm_data_invalid + PatternMatcher([ (UPat.var("x") ^ UPat.var("x"), lambda x: x.const_like(0)), # x^x -> 0 (UPat.var("x") & 0, lambda x: x.const_like(0)), # x&0 -> 0 # (x&mask)>>k -> x>>k when mask only clears bits below k - # TODO: combine this with "# rules for threefry" below ((UPat.var("x") & UPat.cvar("mask")) >> UPat.cvar("k"), lambda x,mask,k: x >> k.val if mask.val | ((1 << k.val) - 1) == -1 else None), ((UPat.var("x") & UPat.cvar("mask")) // UPat.cvar("c"), @@ -168,13 +167,10 @@ symbolic_simple = pm_data_invalid + PatternMatcher([ (UPat.var("x").alu(Ops.POW, UPat.cvar("c")), simplify_pow), # positive const ** x (UPat.cvar("c").alu(Ops.POW, UPat.var("x")), lambda c,x: c if c.val == 1 else (x*math.log2(c.val)).exp2() if c.val > 0 else None), - # rules for threefry - ((UPat.var('x', dtypes.uint64)&0xFFFFFFFF).cast(dtypes.uint32), lambda x: x.cast(dtypes.uint32)), - (((UPat.var(None, dtypes.uint64)*(1<<32)) | UPat.var('y', dtypes.uint32).cast(dtypes.uint64)).cast(dtypes.uint32), lambda y: y), - (((UPat.var('x', dtypes.uint64)*(1<<32)) | UPat.var(None, dtypes.uint32).cast(dtypes.uint64))//(1<<32), lambda x: x), - (((UPat.var(None, dtypes.uint64)<<32) | UPat.var('y', dtypes.uint32).cast(dtypes.uint64)).cast(dtypes.uint32), lambda y: y), - (((UPat.var('x', dtypes.uint64)<<32) | UPat.var(None, dtypes.uint32).cast(dtypes.uint64))//(1<<32), lambda x: x), - (((UPat.var('x', dtypes.uint64)<<32) | UPat.var(None, dtypes.uint32).cast(dtypes.uint64))>>32, lambda x: x), + # unpack a uint64 packed from two uint32 (threefry) + (((UPat.var(None, dtypes.uint64)<<32) | UPat.var('y', dtypes.uint32).cast(dtypes.uint64)).cast(dtypes.uint32), lambda y: y), + (((UPat.var('x', dtypes.uint32).cast(dtypes.uint64)<<32) | UPat.var(None, dtypes.uint32).cast(dtypes.uint64))>>32, + lambda x: x.cast(dtypes.uint64)), # ** simple where folding ** # a conditional with the same results either way is a noop, also fold const conditionals (UPat.var().where(UPat.var("val"), UPat.var("val")), lambda val: val),