add locals

This commit is contained in:
2025-03-28 18:52:48 +08:00
parent 1a9d7a1628
commit e0fd84dd64
2 changed files with 10 additions and 2 deletions
+9 -1
View File
@@ -43,6 +43,8 @@ if __name__ == "__main__":
k.apply_opt(Opt(OptOps.UNROLL, 0, 0))
k.apply_opt(Opt(OptOps.UPCAST, len(k.full_shape)-3, 32))
if k.full_shape[-4]%4 == 0: k.apply_opt(Opt(OptOps.UPCAST, len(k.full_shape)-4, 4))
# if this is small, swap it
if k.full_shape[0] <= 6: k.apply_opt(Opt(OptOps.SWAP, 0, 1))
elif len(k.full_shape) == 3 and k.full_shape[1] == 32:
if k.full_shape[0]%4 != 0: k.apply_opt(Opt(OptOps.PADTO, 0, 4))
# weight without more
@@ -56,7 +58,7 @@ if __name__ == "__main__":
k.apply_opt(Opt(OptOps.UPCAST, 2, 32))
if k.full_shape[1]%4 == 0: k.apply_opt(Opt(OptOps.UPCAST, 1, 4))
# if this is small, just upcast it
if k.full_shape[0] <= 3: k.apply_opt(Opt(OptOps.UPCAST, 0, 0))
if k.full_shape[0] <= 6: k.apply_opt(Opt(OptOps.UPCAST, 0, 0))
elif len(k.full_shape) == 2:
if k.full_shape[0]%128 == 0: k.apply_opt(Opt(OptOps.UPCAST, 0, 128))
elif len(k.full_shape) == 1:
@@ -65,6 +67,12 @@ if __name__ == "__main__":
if k.full_shape[0]%sz == 0:
k.apply_opt(Opt(OptOps.UPCAST, 0, sz))
break
if k.full_shape[0]%2 == 0 and False:
k.apply_opt(Opt(OptOps.LOCAL, 0, k.full_shape[0]//2))
for i in range(1, k.first_reduce-1): k.apply_opt(Opt(OptOps.LOCAL, 1, 0))
else:
# TODO: fix padding
for i in range(1, k.first_reduce): k.apply_opt(Opt(OptOps.LOCAL, 1, 0))
p2 = k.to_program()
new_ei = replace(ei, prg=CompiledRunner(p2), bufs=dsp_bufs)
new_ei.run()
+1 -1
View File
@@ -374,7 +374,7 @@ class Kernel:
if opt.op is OptOps.LOCAL: # cyan
# NOTE: LLVM/CPU can use locals too, but they are treated the same as globals (still helpful for L1 cache)
# it's disabled for now since it makes BEAM slow for little gain
check(self.opts.has_local, "target does not support local")
#check(self.opts.has_local, "target does not support local")
check(axis < self.global_dims, "local is for globals")
self.shift_to(axis, amt, insert_before=self.first_reduce)
self.local_dims += 1