tests passing

This commit is contained in:
2025-10-08 17:06:43 +08:00
parent e3f073b957
commit cd8942cc77
+27 -9
View File
@@ -749,16 +749,17 @@ pm_substitute_recurse = PatternMatcher([(UPat(Ops.SUBSTITUTE, src=(UPat(), UPat(
# *** fast rangeify ***
def apply_rangeify(ctx, x:UOp):
if x.op in {Ops.BUFFERIZE, Ops.INDEX}: return None
if x.op in {Ops.BUFFERIZE, Ops.INDEX, Ops.KERNEL}: return None
if x.op is Ops.ASSIGN and x.src[1].op is Ops.KERNEL: return None
realize_map, range_map, pads_gate = ctx
new_srcs = []
for s in x.src:
new_src = s
if s.op is Ops.BUFFER:
if x in range_map: new_src = new_src.index(*range_map[x][0])
if s.op in {Ops.BUFFER, Ops.MSTACK, Ops.MSELECT} or (s.op is Ops.ASSIGN and s.src[1].op is Ops.KERNEL):
if x in range_map and x.op not in {Ops.MSTACK, Ops.MSELECT}: new_src = new_src.index(*range_map[x][0])
elif s in realize_map:
new_src = UOp(Ops.BUFFERIZE, s.dtype, src=(s,)+tuple(range_map[s][1]), arg=BufferizeOpts(device=s.device), tag=s.tag)
if x in range_map: new_src = new_src.index(*range_map[x][0])
if x in range_map and x.op not in {Ops.MSTACK, Ops.MSELECT}: new_src = new_src.index(*range_map[x][0])
new_srcs.append(new_src)
# NOTE: do we need this?
return x.replace(src=tns) if x.src != (tns:=tuple(new_srcs)) else None
@@ -782,17 +783,29 @@ def remove_movement(ctx, x:UOp):
realize_map, range_map, pads_gate = ctx
if x in range_map or x.src[0].op in {Ops.BUFFERIZE, Ops.INDEX}: return x.src[0]
def fix_assign(ctx, assign:UOp):
realize_map, range_map, pads_gate = ctx
if assign.src[1].op is Ops.KERNEL: return None
to_mop = graph_rewrite(assign.src[0], PatternMatcher([(UPat(GroupOp.Movement, name="x"), lambda x: x.replace(tag=()))]))
ret = assign.replace(src=assign.src+(to_mop,))
range_map[ret] = range_map[assign]
return ret
pm_apply_rangeify = PatternMatcher([
# REDUCE_AXIS -> REDUCE
(UPat(Ops.REDUCE_AXIS, name="x"), fix_reduce_axis),
# PAD -> WHERE
(UPat(Ops.PAD, name="x"), apply_pad),
# add third op to assign
(UPat(Ops.ASSIGN, src=(UPat(), UPat()), name="assign"), fix_assign),
# finally, apply_rangeify
(UPat(GroupOp.All, name="x"), apply_rangeify),
# remove movement op
(UPat(GroupOp.Movement, name="x"), remove_movement),
# const/define_var shouldn't have src
(UPat((Ops.CONST, Ops.DEFINE_VAR), name="c"), lambda ctx,c: c.replace(src=()) if c in ctx[1] else None),
# fixup M index
(UPat((Ops.MSTACK, Ops.MSELECT), name="m"), lambda m: m.replace(src=tuple([x.src[0] if x.op is Ops.INDEX else x for x in m.src]))),
])
@track_rewrites(lambda _,ret: f"Schedule {pluralize('Kernel', len([u for u in UOp.sink(*ret.values()).toposort() if u.op is Ops.KERNEL]))}", True)
@@ -827,11 +840,14 @@ def get_rangeify_map(sink:UOp) -> dict[UOp, UOp]:
# *** these are the ranges on the output ***
consumer_rngs = [range_map[c][0] for c in consumer_map[x] if c.op is not Ops.SINK]
consumer_rngs = [range_map[c][0] for c in consumer_map[x] if c in range_map]
if x in realize_map:
# if this is in the realize_map, we create new ranges (at the output)
#assert x.op not in GroupOp.Movement
out_rngs = [rangeify_ctx.new_range(s) for s in x.shape]
elif x.op is Ops.MSTACK:
# treat MSTACK like SINK
continue
elif len(consumer_rngs) == 0:
continue
elif len(consumer_rngs) == 1:
@@ -861,14 +877,15 @@ def get_rangeify_map(sink:UOp) -> dict[UOp, UOp]:
# we have to realize here if there's new ranges
if not all_all_same: realize_map[x] = None
assert len(out_rngs) == len(x.shape), \
f"shape len mismatch {len(out_rngs)} != {len(x.shape)} on {x.op} with {len(consumer_map[x])} consumers and realize {x in realize_map}"
#assert len(out_rngs) == len(x.shape), \
# f"shape len mismatch {len(out_rngs)} != {len(x.shape)} on {x.op} with {len(consumer_map[x])} consumers and realize {x in realize_map}"
# rngs is the input ranges
rngs = out_rngs[:]
rngs = out_rngs
# handle REDUCE
if x.op is Ops.REDUCE_AXIS:
rngs = rngs[:]
for i,s in enumerate(x.src[0].shape):
if i in x.arg[1]: rngs[i] = rangeify_ctx.new_range(s, axistype=AxisType.REDUCE)
@@ -925,7 +942,8 @@ def get_rangeify_map(sink:UOp) -> dict[UOp, UOp]:
# rebuild the sink with all the BUFFERIZEs with tags, this is what's ending up in the tensor graph
# MSTACK stacks multiple BUFFERIZEs in one tagged tensor
# if it's not tagged by here, it's out
tsink = UOp.sink(*[x for x in tsink.backward_slice if x.base.op in {Ops.BUFFERIZE, Ops.MSTACK, Ops.CONST, Ops.BUFFER} and x.tag is not None])
tsink = UOp.sink(*[x for x in tsink.backward_slice if x.base.op in {Ops.BUFFERIZE, Ops.MSTACK, Ops.CONST, Ops.BUFFER} and \
x.tag is not None and len(x.tag)])
if getenv("VIZ"): graph_rewrite(tsink, PatternMatcher([]), name="View Tagged Rangeify")