diff --git a/extra/replay_pkl.py b/extra/replay_pkl.py index 6e50b8ce63..78aebae441 100644 --- a/extra/replay_pkl.py +++ b/extra/replay_pkl.py @@ -112,6 +112,10 @@ if __name__ == "__main__": k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) k.apply_opt(Opt(OptOps.UPCAST, 1, 96)) k.apply_opt(Opt(OptOps.UPCAST, 0, 4)) + elif knum == 6: + k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) + k.apply_opt(Opt(OptOps.UPCAST, 1, 24)) + k.apply_opt(Opt(OptOps.UPCAST, 0, 16)) elif knum == 37: k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) k.apply_opt(Opt(OptOps.UPCAST, 1, 384)) diff --git a/tinygrad/codegen/symbolic.py b/tinygrad/codegen/symbolic.py index 85ff09be8d..684c0782cb 100644 --- a/tinygrad/codegen/symbolic.py +++ b/tinygrad/codegen/symbolic.py @@ -188,10 +188,10 @@ gep_pushing = PatternMatcher([ # VECTORIZE on same GEP (UPat(Ops.VECTORIZE, name="v", src=UPat(Ops.GEP, src=(UPat.var("x"),))), lambda v,x: x.gep(tuple(get_single_element(i.arg) for i in v.src))), # CAST on multi GEP - (UPat(Ops.CAST, src=(UPat(Ops.GEP, name="g"),), name="c"), - lambda c,g: g.src[0].gep(g.arg[0]).cast(c.dtype.scalar()).broadcast(len(g.arg)) if len(g.arg) > 1 and all_same(g.arg) else None), + #(UPat(Ops.CAST, src=(UPat(Ops.GEP, name="g"),), name="c"), + # lambda c,g: g.src[0].gep(g.arg[0]).cast(c.dtype.scalar()).broadcast(len(g.arg)) if len(g.arg) > 1 and all_same(g.arg) else None), # VECTORIZE/CONST - (UPat(Ops.VECTORIZE, src=UPat.var("x"))+UPat.cvar("c", vec=False), lambda x,c: (x+c.arg).broadcast(c.dtype.count)), + #(UPat(Ops.VECTORIZE, src=UPat.var("x"))+UPat.cvar("c", vec=False), lambda x,c: (x+c.arg).broadcast(c.dtype.count)), ]) symbolic = symbolic_simple+PatternMatcher([ diff --git a/tinygrad/runtime/ops_dsp.py b/tinygrad/runtime/ops_dsp.py index d9cbf3562d..efb3d0fc0c 100644 --- a/tinygrad/runtime/ops_dsp.py +++ b/tinygrad/runtime/ops_dsp.py @@ -21,7 +21,6 @@ def multi_mul(a0, a1, b0, b1, c0, c1, d0, d1, acc=None): swizzle.append(96+i) swizzle = tuple(swizzle) if a0.op is not Ops.CAST: return None - if d0.op is not Ops.CAST: return None if a1.op is not Ops.CAST: return None assert a0.op is Ops.CAST assert b0.op is Ops.CAST