From 25c023bcbe3beb07f9f8bd13187fe57bacaa6126 Mon Sep 17 00:00:00 2001 From: George Hotz Date: Sat, 22 Mar 2025 17:49:34 +0800 Subject: [PATCH] more --- extra/replay_pkl.py | 25 +++++++++++++++++++++++-- 1 file changed, 23 insertions(+), 2 deletions(-) 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]))