FLImaging 7.8.25.3
BackendSSIM.h
1#pragma once
2
3#if _MSC_VER >= 1900 && defined(_M_X64)
4
5#include "BackendBase.h"
6#include "BackendConv2D.h"
7
8namespace FLImaging
9{
10 namespace AI
11 {
12 template <typename T>
13 class FL_EXPORT CTensor;
14
15 template <typename T>
16 class FL_EXPORT CBackendSSIM : public CBackendBase<T>
17 {
18 public:
19 CBackendSSIM();
20 CBackendSSIM(const CBackendSSIM<T>& bs);
21 virtual ~CBackendSSIM();
22
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);
25
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);
28
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;
31
32 virtual const int64_t GetRequiredDedicatedMemory(std::vector<int64_t>& vctInputShape, int64_t i64Channel) const;
33
34 DeclareGetClassType();
35 SupportToDuplicateObjectWithoutCreateNewObject(CBackendSSIM<T>, *this);
36
37 protected:
38 T m_tMaxDiff;
39 int64_t m_i64FilterSize;
40 T m_tFilterSigma;
41 T m_tK1;
42 T m_tK2;
43 bool m_bReturnIndexMap;
44
45 CTensor<T> m_tsrGaussianFilter;
46
47 CTensor<T> m_tsrMeanX;
48 CTensor<T> m_tsrMeanY;
49 CTensor<T> m_tsrMap;
50 CTensor<T> m_tsrA;
51 CTensor<T> m_tsrB;
52 CTensor<T> m_tsrC;
53 CTensor<T> m_tsrD;
54
55 protected:
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);
60 };
61 }
62}
63
64#endif
Definition AlgorithmAIBase.h:18