forked from tinygrad/tinygrad
remove noop from to_movement_ops (#2693)
This commit is contained in:
+37
-24
@@ -1,12 +1,31 @@
|
||||
from tqdm import tqdm
|
||||
import itertools
|
||||
from typing import List, Tuple
|
||||
from collections import defaultdict
|
||||
from typing import List, Tuple, DefaultDict
|
||||
from extra.optimization.helpers import load_worlds, ast_str_to_ast
|
||||
from tinygrad.ops import MovementOps, BufferOps, LazyOp
|
||||
from tinygrad.helpers import prod
|
||||
from tinygrad.shape.shapetracker import ShapeTracker
|
||||
from tinygrad.shape.symbolic import sym_infer, Node
|
||||
|
||||
def apply_mop(st: ShapeTracker, mop_arg: Tuple[MovementOps, Tuple]) -> ShapeTracker:
|
||||
mop, arg = mop_arg
|
||||
if mop == MovementOps.RESHAPE:
|
||||
# shapetracker doesn't allow flattening with -1 but required for MovementOps.RESHAPE
|
||||
if arg == (-1,): return st.reshape((prod(st.views[-1].shape),))
|
||||
return st.reshape(arg)
|
||||
if mop == MovementOps.PERMUTE: return st.permute(arg)
|
||||
if mop == MovementOps.EXPAND:
|
||||
if len(arg) != len(st.shape): st = st.reshape((1,*st.shape))
|
||||
return st.expand(arg)
|
||||
if mop == MovementOps.PAD: return st.pad(arg)
|
||||
if mop == MovementOps.SHRINK: return st.shrink(arg)
|
||||
if mop == MovementOps.STRIDE: return st.stride(arg)
|
||||
raise ValueError("invalid mop")
|
||||
|
||||
def make_scratch_st(st: ShapeTracker) -> ShapeTracker:
|
||||
return ShapeTracker.from_shape((get_buffer_size(st.views[0].shape, st.views[0].strides, st.views[0].offset, st.views[0].mask),))
|
||||
|
||||
# ShapeTracker to an equivalent series of MovementOps (https://github.com/tinygrad/tinygrad/pull/2216)
|
||||
def to_movement_ops(st: ShapeTracker) -> List[Tuple[MovementOps, Tuple]]:
|
||||
to_apply:List[Tuple[MovementOps, Tuple]] = []
|
||||
@@ -50,7 +69,16 @@ def to_movement_ops(st: ShapeTracker) -> List[Tuple[MovementOps, Tuple]]:
|
||||
if any(s != 1 and st == 0 for s,st in zip(real_shape, v.strides)): to_apply.append((MovementOps.EXPAND, real_shape))
|
||||
# lastly, we apply post expand pads
|
||||
if v.mask is not None and any(x != (0,0) for x in post_expand_pads): to_apply.append((MovementOps.PAD, post_expand_pads))
|
||||
return to_apply
|
||||
|
||||
scratch_st = make_scratch_st(st)
|
||||
ret = []
|
||||
for mop_arg in to_apply:
|
||||
st = apply_mop(scratch_st, mop_arg)
|
||||
if st != scratch_st:
|
||||
ret.append(mop_arg)
|
||||
scratch_st = st
|
||||
|
||||
return ret
|
||||
|
||||
def get_real_view(shape, strides, offset, mask):
|
||||
real_shape = tuple(y-x for x,y in mask) if mask else shape
|
||||
@@ -90,29 +118,12 @@ def st_equivalent(st1: ShapeTracker, st2: ShapeTracker):
|
||||
|
||||
return True
|
||||
|
||||
c: DefaultDict[int,int] = defaultdict(int)
|
||||
def test_rebuild(st: ShapeTracker):
|
||||
rebuilt_st = ShapeTracker.from_shape((get_buffer_size(st.views[0].shape, st.views[0].strides, st.views[0].offset, st.views[0].mask),))
|
||||
for mop, arg in to_movement_ops(st):
|
||||
if mop == MovementOps.RESHAPE:
|
||||
# shapetracker doesn't allow flattening with -1 but required for MovementOps.RESHAPE
|
||||
if arg == (-1,):
|
||||
rebuilt_st = rebuilt_st.reshape((prod(rebuilt_st.views[-1].shape),))
|
||||
else:
|
||||
rebuilt_st = rebuilt_st.reshape(arg)
|
||||
elif mop == MovementOps.PERMUTE:
|
||||
rebuilt_st = rebuilt_st.permute(arg)
|
||||
elif mop == MovementOps.EXPAND:
|
||||
if len(arg) != len(rebuilt_st.shape):
|
||||
rebuilt_st = rebuilt_st.reshape((1,*rebuilt_st.shape))
|
||||
rebuilt_st = rebuilt_st.expand(arg)
|
||||
elif mop == MovementOps.PAD:
|
||||
rebuilt_st = rebuilt_st.pad(arg)
|
||||
elif mop == MovementOps.SHRINK:
|
||||
rebuilt_st = rebuilt_st.shrink(arg)
|
||||
elif mop == MovementOps.STRIDE:
|
||||
rebuilt_st = rebuilt_st.stride(arg)
|
||||
else:
|
||||
raise Exception("invalid mop")
|
||||
rebuilt_st = make_scratch_st(st)
|
||||
mops = to_movement_ops(st)
|
||||
c[len(mops)] += 1
|
||||
for mop_arg in mops: rebuilt_st = apply_mop(rebuilt_st, mop_arg)
|
||||
rebuilt_st = rebuilt_st.simplify()
|
||||
assert st_equivalent(st, rebuilt_st)
|
||||
last_v1 = st.views[-1]
|
||||
@@ -130,3 +141,5 @@ if __name__ == "__main__":
|
||||
ast_strs = load_worlds(False, False, True)[:4000]
|
||||
for ast_str in tqdm(ast_strs):
|
||||
test_interpret_ast(ast_str_to_ast(ast_str))
|
||||
|
||||
print(f"avg length of mop = {sum(k*v for k,v in c.items()) / sum(c.values()):.2f}")
|
||||
Reference in New Issue
Block a user