forked from tinygrad/tinygrad
more
This commit is contained in:
+23
-2
@@ -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]))
|
||||
|
||||
Reference in New Issue
Block a user