This commit is contained in:
2025-03-22 17:49:34 +08:00
parent 07abf9e6bc
commit 25c023bcbe
+23 -2
View File
@@ -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]))