diff --git a/extra/onnx.py b/extra/onnx.py index f2862a9e25..d4c77d7c3d 100644 --- a/extra/onnx.py +++ b/extra/onnx.py @@ -764,7 +764,11 @@ def get_onnx_ops(): else: # DSP swizzle memory x = x.reshape(x.shape[0], x.shape[1]//WEIGHT_SHIFT, WEIGHT_SHIFT).permute(1,0,2).contiguous().permute(1,0,2).reshape(x.shape) - #if getenv("NHWC") and x.shape == (1000, 1280): + if getenv("NHWC") and x.shape == (1000, 1280): + x = x.reshape(-1, 320, 4) + order = (1,0,2) + x = x.permute(*order).contiguous().permute(*argsort(order)) + x = x.reshape(-1, 1280) x_scale, x_zero_point = _prepare_quantize(x, x_scale, x_zero_point, axis, block_size) return ((x.int() - x_zero_point) * x_scale).cast(x_scale.dtype) diff --git a/extra/replay_pkl.py b/extra/replay_pkl.py index 3c2fb85be4..9245c5d3fc 100644 --- a/extra/replay_pkl.py +++ b/extra/replay_pkl.py @@ -118,23 +118,24 @@ if __name__ == "__main__": k.apply_opt(Opt(OptOps.UNROLL, 0, 8)) k.apply_opt(Opt(OptOps.UPCAST, 1, 96)) k.apply_opt(Opt(OptOps.UPCAST, 0, 4)) + elif knum == 26: + # 14x14, 384 chan, 3x3 dwconv + k.apply_opt(Opt(OptOps.UNROLL, 1, 0)) + k.apply_opt(Opt(OptOps.UPCAST, 2, 128)) elif knum == 5: k.apply_opt(Opt(OptOps.UNROLL, 1, 0)) k.apply_opt(Opt(OptOps.UPCAST, 2, 0)) k.apply_opt(Opt(OptOps.UPCAST, 1, 4)) # this breaks something #k.apply_opt(Opt(OptOps.UPCAST, 1, 4)) - elif knum == 8: - # 3x3 dwconv w 144 chans on 56x56 - k.apply_opt(Opt(OptOps.UNROLL, 1, 0)) - k.apply_opt(Opt(OptOps.UPCAST, 2, 0)) + #elif knum in [8, 12]: + # 3x3 dwconv w 144 chans on 56x56 / 28x28 + #k.apply_opt(Opt(OptOps.UNROLL, 1, 0)) + #k.apply_opt(Opt(OptOps.UPCAST, 2, 0)) #k.apply_opt(Opt(OptOps.UPCAST, 1, 4)) - elif knum in [8, 12]: - # 3x3 dwconv w 144 chans - k.apply_opt(Opt(OptOps.UPCAST, 2, 144)) - elif knum in [15, 19]: + #elif knum in [15, 19]: # 3x3 dwconv w 192 chans - k.apply_opt(Opt(OptOps.UPCAST, 2, 192)) + #k.apply_opt(Opt(OptOps.UPCAST, 2, 192)) elif knum == 6: k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) k.apply_opt(Opt(OptOps.UPCAST, 1, 24)) @@ -185,15 +186,27 @@ if __name__ == "__main__": elif knum == 66: k.apply_opt(Opt(OptOps.UNROLL, 0, 4)) k.apply_opt(Opt(OptOps.UPCAST, 0, 8)) + #k.apply_opt(Opt(OptOps.UPCAST, 0, 8)) #k.apply_opt(Opt(OptOps.PADTO, 0, 32)) #k.apply_opt(Opt(OptOps.UPCAST, 0, 32)) + pass else: full_shape = k.full_shape out_shape = k.sts[0].shape out_strides = k.sts[0].real_strides() - # there's some bug here with this - #if len(out_strides) == 5 and full_shape[-2:] == (3,3): - #if knum not in [2,30]: k.apply_opt(Opt(OptOps.UPCAST, 2, 0)) + if len(out_strides) == 5 and full_shape[-2:] == (3,3): + # 3x3 dwconv + k.apply_opt(Opt(OptOps.UNROLL, 1, 0)) + if full_shape[2]%128 == 0: + # optimal + k.apply_opt(Opt(OptOps.UPCAST, 2, 128)) + elif full_shape[2]%64 == 0: + # sub-optimal 64 + k.apply_opt(Opt(OptOps.UPCAST, 2, 64)) + elif full_shape[2] == 144: + # bad 144 + k.apply_opt(Opt(OptOps.UPCAST, 2, 144)) + else: raise RuntimeError(f"3x3 conv missing {full_shape}") if len(out_strides) == 3: if full_shape[1] == 192 and full_shape[0]%2 == 0: k.apply_opt(Opt(OptOps.UNROLL, 0, 4))