// // MetalLayerNorm.hpp // MNN // // Created by MNN on 2019/01/30. // Copyright © 2018, Alibaba Group Holding Limited // #ifndef MetalLayerNorm_hpp #define MetalLayerNorm_hpp #import "MetalExecution.hpp" #import "MNN_generated.h" #if MNN_METAL_ENABLED namespace MNN { class MetalLayerNorm : public MetalExecution { public: struct Resource { int mGroup = 1; float mEps; int mAxisSize; bool mHasGammaBeta = false; bool mRMSNorm = false; int mGammaSize = 0; std::shared_ptr mGammaBuffer; std::shared_ptr mBetaBuffer; }; MetalLayerNorm(Backend *backend, std::shared_ptr res); virtual ~MetalLayerNorm() = default; virtual ErrorCode onResize(const std::vector &inputs, const std::vector &outputs) override; virtual void onEncode(const std::vector &inputs, const std::vector &outputs, id encoder) override; static std::shared_ptr makeResource(Backend *backend, const LayerNorm *layernorm); virtual bool onClone(Backend* bn, const Op* op, Execution** dst) override; private: int mOutside; int mInside; bool mIsNC4HW4 = false; bool mIsBinaryNCHW = false; int mChannelUnit; std::shared_ptr mResource; id mShapeBuffer; id mPipeline; std::pair mThreads; }; } // namespace MNN #endif /* MNN_METAL_ENABLED */ #endif /* MetalLayerNorm_hpp */