// // MetalConvolutionWinograd.hpp // MNN // // Created by MNN on 2019/01/31. // Copyright © 2018, Alibaba Group Holding Limited // #ifndef MetalConvolutionWinograd_hpp #define MetalConvolutionWinograd_hpp #import "MetalConvolutionCommon.hpp" #if MNN_METAL_ENABLED namespace MNN { struct TransformBuffer { int inputSize[4]; int outputSize[4]; int padX; int padY; int unitWidth; int unitHeight; int unit; int activation; int remain[2]; }; class MetalConvolutionWinograd : public MetalConvolutionCommon { public: static bool isValid(Backend *backend, const Convolution2D *conv, const Tensor *input, const Tensor* output); MetalConvolutionWinograd(Backend *backend, const MNN::Op *op); virtual ~MetalConvolutionWinograd() = default; virtual ErrorCode onResize(const std::vector &inputs, const std::vector &outputs) override; virtual bool onClone(Backend* bn, const Op* op, Execution** dst) override; protected: virtual void onEncode(const std::vector &inputs, const std::vector &outputs, id encoder) override; virtual std::shared_ptr weightTransform(int group, int oc, int ic, int kh, int kw, const float *src, bool int8Weight=false, bool int4Weight=false, id srcGpuBuffer=nil, int subBits=0) override; private: MetalConvolutionWinograd(Backend *backend, const MNN::Op *op, std::shared_ptr weight, std::shared_ptr bias); id mShapeBuffer = nil; int mSrcUnit; int mDstUnit; std::shared_ptr mTempSrc; std::shared_ptr mTempDst; MTLSize mInputTransformThreads; MTLSize mMatMulThreads; MTLSize mOutputTransformThreads; int mSplitNum = 1; }; } // namespace MNN #endif /* MNN_METAL_ENABLED */ #endif /* MetalConvolutionWinograd_hpp */