diff --git a/extra/replay_pkl.py b/extra/replay_pkl.py index f0f0c46061..e7ff7bc764 100644 --- a/extra/replay_pkl.py +++ b/extra/replay_pkl.py @@ -116,7 +116,7 @@ if __name__ == "__main__": 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 == 11: + elif knum in [7,11]: k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) k.apply_opt(Opt(OptOps.UPCAST, 1, 144)) #k.apply_opt(Opt(OptOps.UPCAST, 0, 8)) @@ -124,11 +124,24 @@ if __name__ == "__main__": k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) k.apply_opt(Opt(OptOps.UPCAST, 1, 192)) k.apply_opt(Opt(OptOps.UPCAST, 0, 2)) + elif knum == 40: + k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) + k.apply_opt(Opt(OptOps.UPCAST, 1, 64)) + #k.apply_opt(Opt(OptOps.UPCAST, 0, 2)) + pass + #elif knum == 18: + # k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) + # k.apply_opt(Opt(OptOps.UPCAST, 1, 192)) + # k.apply_opt(Opt(OptOps.UPCAST, 0, 2)) #elif knum == 33: # 196x64 * 64x384 -> 196x384 # automatic gets this now #k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) #k.apply_opt(Opt(OptOps.UPCAST, 1, 128)) + #elif knum == 39: + #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 == 37: k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) k.apply_opt(Opt(OptOps.UPCAST, 1, 384)) @@ -137,7 +150,15 @@ if __name__ == "__main__": out_shape = k.sts[0].shape out_strides = k.sts[0].real_strides() if len(out_strides) == 3: - if full_shape[1] < 128: + if full_shape[1] == 192 and full_shape[0]%2 == 0: + k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) + k.apply_opt(Opt(OptOps.UPCAST, 1, 192)) + k.apply_opt(Opt(OptOps.UPCAST, 0, 2)) + elif full_shape[1] == 96 and full_shape[0]%4 == 0: + 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 full_shape[1] < 128: if full_shape[2] <= 16: k.apply_opt(Opt(OptOps.UNROLL, 0, 0)) else: k.apply_opt(Opt(OptOps.UNROLL, 0, 8)) k.apply_opt(Opt(OptOps.UPCAST, 1, full_shape[1]))