diff --git a/test/amd/test_asm_kernel.py b/test/amd/test_asm_kernel.py index 5ba12765cd..c7b6b036c4 100644 --- a/test/amd/test_asm_kernel.py +++ b/test/amd/test_asm_kernel.py @@ -36,7 +36,7 @@ def custom_add_var(A:UOp, B:UOp) -> UOp: A,B = A.flatten(), B.flatten() assert A.dtype == dtypes.uint32, f"buffer dtype must be uint32, got {A.dtype}" threads = UOp.special(A.numel(), "lidx0") - var = UOp.param(2, dtypes.weakint, vmin_vmax=(0, 10), name="var", addrspace=AddrSpace.ALU) + var = UOp.param(2, dtypes.int, vmin_vmax=(0, 10), name="var", addrspace=AddrSpace.ALU) insts = [ s_load_b128(s[4:7], s[0:1]), s_load_b32(s[8], s[0:1], offset=0x10), # all threads load the same variable diff --git a/test/null/test_uop_symbolic.py b/test/null/test_uop_symbolic.py index bc9e39d65b..e7e5546214 100644 --- a/test/null/test_uop_symbolic.py +++ b/test/null/test_uop_symbolic.py @@ -987,7 +987,7 @@ class TestSymbolic(unittest.TestCase): self.helper_test_variable(cond.ne(False), 0, 1, "(x<2)") def test_bitcast_chain(self): - a = Variable("a", 0, 3) + a = UOp.variable("a", 0, 3, dtype=dtypes.int32) self.assertIs(graph_rewrite(a.bitcast(dtypes.float32).bitcast(a.dtype), sym), a) def test_negation_in_where(self): diff --git a/tinygrad/callify.py b/tinygrad/callify.py index 734fc68c2c..0fce987b0c 100644 --- a/tinygrad/callify.py +++ b/tinygrad/callify.py @@ -179,10 +179,9 @@ def finalize_after(ctx:AllocCtx, x:UOp): def replace_input_buffer(ctx:AllocCtx, b:UOp): ctx.replacements.append(b) + if b.op is Ops.BIND: return b.param_like(len(ctx.replacements)-1) return UOp.param(len(ctx.replacements)-1, b.dtype, b.shape, b.device, - b._min_max if b.op is Ops.BIND else None, name=b.src[0].expr if b.op is Ops.BIND else None, - addrspace=b.addrspace if b.addrspace is not None else AddrSpace.GLOBAL, - multiple_of=b.src[0].arg.multiple_of if b.op is Ops.BIND else None) + addrspace=b.addrspace if b.addrspace is not None else AddrSpace.GLOBAL) pm_finalize_call = PatternMatcher([ (UPat(Ops.AFTER, name="x"), finalize_after), diff --git a/tinygrad/dtype.py b/tinygrad/dtype.py index 6defe35113..513515dba9 100644 --- a/tinygrad/dtype.py +++ b/tinygrad/dtype.py @@ -153,7 +153,7 @@ class dtypes: uints = (uint8, uint16, uint32, uint64) sints = (int8, int16, int32, int64) ints = uints + sints - weaks = (weakfloat,) + weaks = (weakint, weakfloat) all = floats + ints + (bool,) # noqa: A003 if (env_default_float := getenv("DEFAULT_FLOAT", "")):