diff --git a/tinygrad/codegen/__init__.py b/tinygrad/codegen/__init__.py index 17693ed9d4..39eb47c359 100644 --- a/tinygrad/codegen/__init__.py +++ b/tinygrad/codegen/__init__.py @@ -135,7 +135,7 @@ def do_stack_wmma(u:UOp): src = [] for b in u.src: if b.op != Ops.STACK: - src.append(UOp.stack(*[b.index(UOp.const(None, i)) for i in range(b.max_numel())])) + src.append(UOp.stack(*[b.index(i) for i in range(b.max_numel())])) else: src.append(b) return u.replace(src=tuple(src)) @@ -161,7 +161,7 @@ devectorizer2 = mop_cleanup+pm_mops+PatternMatcher([ # RESHAPE a void is removed (hack for AFTER) (UPat(Ops.RESHAPE, dtype=dtypes.void, name="x"), lambda x: x.src[0]), # reshape of a single element shaped value to scalar is an index - (UPat(Ops.RESHAPE, name="x"), lambda x: x.src[0].index(UOp.const(None, 0)) if x.marg == () and x.src[0].shape == (1,) else None), + (UPat(Ops.RESHAPE, name="x"), lambda x: x.src[0].index(0) if x.marg == () and x.src[0].shape == (1,) else None), # EXPAND on scalar -> STACK (UPat(Ops.EXPAND, src=(UPat.var("x"), UPat()), name="out"), lambda x,out: UOp.stack(*([x]*out.max_numel())) if x.shape == () and out.shape == (out.max_numel(),) else None), diff --git a/tinygrad/codegen/late/coalesce.py b/tinygrad/codegen/late/coalesce.py index 8830c643b4..8510f639bd 100644 --- a/tinygrad/codegen/late/coalesce.py +++ b/tinygrad/codegen/late/coalesce.py @@ -158,7 +158,7 @@ def memory_coalescing(sink:UOp, ctx:Renderer) -> UOp: ld = idx.load() for i,g in enumerate(grp): for oo in offsets[g]: - replacements[oo] = ld.index(UOp.const(None, i)) if len(grp) > 1 else ld + replacements[oo] = ld.index(i) if len(grp) > 1 else ld full_grp = full_grp[length:] # apply diff --git a/tinygrad/llm/gguf.py b/tinygrad/llm/gguf.py index de1482515e..6c486e29fc 100644 --- a/tinygrad/llm/gguf.py +++ b/tinygrad/llm/gguf.py @@ -105,7 +105,7 @@ def ggml_data_to_tensor(t: Tensor, n: int, ggml_type: int) -> Tensor: if ggml_type == 39: e = blocks[:, 0].cast(dtypes.uint32) small_bits = Tensor([0x00200000, 0x00400000], dtype=dtypes.uint32, device=t.device)[e.clip(0, 1).cast(dtypes.int32)] # e = 0 or e = 1 case - d = (e < 2).where(small_bits, ((e - 1) * 0x00800000).cast(dtypes.uint32)).bitcast(dtypes.float32).unsqueeze(-1) + d = (e < 2).where(small_bits, (e - 1) * 0x00800000).bitcast(dtypes.float32).unsqueeze(-1) codes = q_to_uint8(blocks[:, 1:17], 4) fp4_lut = Tensor([0.0, 1.0, 2.0, 3.0, 4.0, 6.0, 8.0, 12.0, -0.0,-1.0,-2.0,-3.0,-4.0,-6.0,-8.0,-12.0], diff --git a/tinygrad/mixin/op.py b/tinygrad/mixin/op.py index 042754cb6a..2b8fc1e17d 100644 --- a/tinygrad/mixin/op.py +++ b/tinygrad/mixin/op.py @@ -139,13 +139,13 @@ class OpMixin(ElementwiseMixin, ReduceMixin): if v is None: return x # advanced getitem # advanced setitem: resolve tensor dims in collapsed space, then fall through to basic setitem path - vb = v.cast(self.dtype)._broadcast_to(_broadcast_shape(x.shape, v.shape)) + vb = v._broadcast_to(_broadcast_shape(x.shape, v.shape)) for dim in sum_axis: vb = vb.unsqueeze(dim) # add back reduced dims from sum start = dims[0] if not permuted else 0 vb = x_pre._masked_merge(vb, mask, tuple(range(start, start + len(big_shape)))) elif v is None: return x # basic getitem # basic setitem: broadcast v, reshape to self.ndim (unsqueeze int dims, squeeze None dims) - else: vb = v.cast(self.dtype)._broadcast_to(x.shape) + else: vb = v._broadcast_to(x.shape) vb = vb.reshape(tuple(1 if p['collapse_dim'] else p['size'] for p in indices_parsed if p['index'] is not None)) per_dim = [] for d, m in enumerate(mops): diff --git a/tinygrad/schedule/rangeify.py b/tinygrad/schedule/rangeify.py index 38e2d36b0f..fdd0601072 100644 --- a/tinygrad/schedule/rangeify.py +++ b/tinygrad/schedule/rangeify.py @@ -78,7 +78,7 @@ def split_reduceop(reduce:UOp, x:UOp): # split is moved to the end to provide maximum locality for the second phase reduce. # get expanded by rangeifying the UOp x - indexed = x.index(*[UOp.range(s, i) if resolve(s>1) else UOp.const(None, 0) for i,s in enumerate(x.shape)]) + indexed = x.index(*[UOp.range(s, i) if resolve(s>1) else 0 for i,s in enumerate(x.shape)]) range_nums = [y.arg[0] for y in indexed.substitute({x.base:UOp(Ops.NOOP, x.base.dtype)}, extra_pm=pm_mops).ranges] is_expanded = [i not in range_nums for i in range(len(x.shape))] diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index 9ebcda3f11..12cd821f81 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -98,7 +98,7 @@ def shape_to_shape_arg(arg:tuple[sint, ...]) -> UOp: if isinstance(x, UOp) and not dtypes.is_int(x.dtype): raise RuntimeError(f"shape must be int, got {x.dtype} in {arg}") if len(arg) == 0: return UOp(Ops.STACK) elif len(arg) == 1: return UOp.const(dtypes.weakint, arg[0]) - else: return UOp(Ops.STACK, src=tuple(UOp.const(dtypes.weakint, x) if isinstance(x, int) else x for x in arg)) + else: return UOp(Ops.STACK, src=tuple(UOp.const(None, x) if isinstance(x, int) else x for x in arg)) def consumer_map_from_toposort(lst:Iterable[UOp]): ret: dict[UOp, dict[UOp, None]] = {} @@ -557,7 +557,7 @@ class UOp(RandMixin, metaclass=UOpMetaClass): if len(srcs) == 1 and isinstance(srcs[0], UOp): return srcs[0] return UOp(Ops.GROUP, src=tuple([x for x in srcs if x is not None])) def index(self, *srcs:UOp|int|None, **kwargs): - new_srcs: list[UOp] = [UOp.const(dtypes.weakint, x) if isinstance(x, int) else x for x in srcs if x is not None] + new_srcs: list[UOp] = [UOp.const(None, x) if isinstance(x, int) else x for x in srcs if x is not None] if len(new_srcs) == 1 and new_srcs[0].op is Ops.CONST and self.op is Ops.STACK: return self.src[new_srcs[0].arg] return UOp(Ops.INDEX, src=(self,)+tuple(new_srcs), **kwargs) def __getitem__(self, idx): @@ -569,11 +569,11 @@ class UOp(RandMixin, metaclass=UOpMetaClass): bounds = tuple((s.start or 0, s.stop if s.stop is not None else self.shape[i]) if isinstance(s, slice) else (0, self.shape[i]) for i, s in enumerate(idx)) src = self.shrink(bounds) - non_slice_args = [UOp.const(dtypes.weakint, x) if isinstance(x, int) else x for x in idx if not isinstance(x, slice)] + non_slice_args = [x for x in idx if not isinstance(x, slice)] if not non_slice_args: return src # all dims are slices, no indexing needed perm = src.permute(tuple([i for i in range(src.ndim) if i not in slice_idx] + slice_idx)) return perm.index(*non_slice_args) - return self.index(*[UOp.const(dtypes.weakint, x) if isinstance(x, int) else x for x in idx]) + return self.index(*idx) @property def _uop(self) -> UOp: return self @classmethod diff --git a/tinygrad/uop/symbolic.py b/tinygrad/uop/symbolic.py index 1d9a7c658d..87d6130c51 100644 --- a/tinygrad/uop/symbolic.py +++ b/tinygrad/uop/symbolic.py @@ -70,16 +70,16 @@ invalid_pat = UPat(Ops.CONST, arg=Invalid, name="i") invalid_gate = UPat.var("cond").where(UPat.var("x"), invalid_pat) pm_data_invalid = PatternMatcher([ (invalid_pat.broadcast(), lambda i: i), - (UPat(GroupOp.Unary|{Ops.BITCAST}, src=(invalid_pat,), name="op"), lambda i,op: i.cast(op.dtype)), + (UPat(GroupOp.Unary|{Ops.BITCAST}, src=(invalid_pat,)), lambda i: i), (UPat(GroupOp.Unary|{Ops.CAST, Ops.BITCAST}, src=(invalid_gate,), name="op"), - lambda cond,x,op,i: cond.where(op.replace(src=(x,)), i.cast(op.dtype))), + lambda cond,x,op,i: cond.where(op.replace(src=(x,)), i)), # binary ops move inside the gate, with Invalid in the false branch - (UPat(GroupOp.Binary, src=(invalid_gate, UPat.var("y")), name="alu"), lambda cond,x,y,alu,i: cond.where(x.alu(alu.op,y), i.cast(alu.dtype))), - (UPat(GroupOp.Binary, src=(UPat.var("y"), invalid_gate), name="alu"), lambda cond,x,y,alu,i: cond.where(y.alu(alu.op,x), i.cast(alu.dtype))), + (UPat(GroupOp.Binary, src=(invalid_gate, UPat.var("y")), name="alu"), lambda cond,x,y,alu,i: cond.where(x.alu(alu.op,y), i)), + (UPat(GroupOp.Binary, src=(UPat.var("y"), invalid_gate), name="alu"), lambda cond,x,y,alu,i: cond.where(y.alu(alu.op,x), i)), (UPat(GroupOp.Binary-GroupOp.Comparison, src=[invalid_pat, UPat()]), lambda i: i), # an Invalid condition poisons the whole where; a gated Invalid condition lifts the gate out - (invalid_pat.where(UPat.var("a"), UPat()), lambda i,a: i.cast(a.dtype)), - (invalid_gate.where(UPat.var("a"), UPat.var("b")), lambda cond,x,i,a,b: cond.where(x.where(a,b), i.cast(a.dtype))), + (invalid_pat.where(UPat(), UPat()), lambda i: i), + (invalid_gate.where(UPat.var("a"), UPat.var("b")), lambda cond,x,i,a,b: cond.where(x.where(a,b), i)), # normalize where(cond, Invalid, val) -> where(~cond, val, Invalid) (UPat.var("cond").where(invalid_pat, UPat.var("val")), lambda cond, i, val: cond.logical_not().where(val, i) if val.arg != Invalid else i), # lift Invalid out: a.where(cond.where(x, Invalid), c) -> (~a|cond).where(a.where(x, c), Invalid)