3#if _MSC_VER >= 1900 && defined(_M_X64)
5#include "BackendBase.h"
6#include "BackendConv2D.h"
13 class FL_EXPORT CTensor;
16 class FL_EXPORT CBackendSSIM :
public CBackendBase<T>
20 CBackendSSIM(
const CBackendSSIM<T>& bs);
21 virtual ~CBackendSSIM();
23 virtual const CResult SetSSIMParameters(T tMaxDiffValue = (T)INFINITY, int64_t i64FilterSize = 11, T tFilterSigma = (T)1.5, T tK1 = (T)0.01, T tK2 = (T)0.03,
bool bReturnIndexMap =
false);
24 virtual const CResult GetSSIMParameters(T& tMaxDiffValue, int64_t& i64FilterSize, T& tFilterSigma, T& tK1, T& tK2,
bool& bReturnIndexMap);
26 virtual const CResult SSIM(
const CTensor<T>* pTsrInput1,
const CTensor<T>* pTsrInput2, CTensor<T>* pTsrResult, CBackendConv2D<T>& backendConv2D,
bool bTraining);
27 virtual const CResult Backward(
const CTensor<T>* pTsrInput1,
const CTensor<T>* pTsrInput2,
const CTensor<T>* pTsrDy, CTensor<T>* pTsrDx1, CTensor<T>* pTsrDx2, CBackendConv2D<T>& backendConv2D);
29 virtual const int64_t GetRequiredTemporaryMemory_Forward(std::vector<int64_t>& vctInputShape)
const;
30 virtual const int64_t GetRequiredTemporaryMemory_Backward(std::vector<int64_t>& vctInputShape)
const;
32 virtual const int64_t GetRequiredDedicatedMemory(std::vector<int64_t>& vctInputShape, int64_t i64Channel)
const;
34 DeclareGetClassType();
35 SupportToDuplicateObjectWithoutCreateNewObject(CBackendSSIM<T>, *
this);
39 int64_t m_i64FilterSize;
43 bool m_bReturnIndexMap;
45 CTensor<T> m_tsrGaussianFilter;
47 CTensor<T> m_tsrMeanX;
48 CTensor<T> m_tsrMeanY;
56 virtual const CResult CreateGaussianFilter(int64_t i64Channel);
57 virtual const CResult Calculate_Forward();
58 virtual const CResult Calculate_Backward(
const CTensor<T>* pTsrDy, CTensor<T>* pTsrdB, CTensor<T>* pTsrdD, T tDivide);
59 virtual const CResult Calculate_UpdateGradient(
const CTensor<T>* pTsrInput1,
const CTensor<T>* pTsrInput2,
const CTensor<T>* pTsrDy, CTensor<T>* pTsrdB, CTensor<T>* pTsrdD, CTensor<T>* pTsrConvB, CTensor<T>* pTsrConvD, CTensor<T>* pTsrConvMu, CTensor<T>* pTsrdMu, CTensor<T>* pTsrDx, CBackendConv2D<T>& backendConv2D, T tDivide,
bool bOpr1);
Definition AlgorithmAIBase.h:18