forked from tinygrad/tinygrad
tests passing
This commit is contained in:
@@ -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")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user