remove no-op explicit dtype= or cast [PR] (#17276)

This commit is contained in:
chenyu
2026-07-29 00:47:35 -04:00
committed by GitHub
parent dd16d5aead
commit 2f8f2d2d37
7 changed files with 17 additions and 17 deletions
+2 -2
View File
@@ -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),
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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],
+2 -2
View File
@@ -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):
+1 -1
View File
@@ -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))]
+4 -4
View File
@@ -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
+6 -6
View File
@@ -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)