FLImaging 7.8.25.3
ComputationalGraphTransConv2D.h
1#pragma once
2
3#if _MSC_VER >= 1900 && defined(_M_X64)
4
5#include "ComputationalGraph.h"
6#include "Tensor.h"
7#include "BackendConv2D.h"
8
9#include <vector>
10
11namespace FLImaging
12{
13 namespace AI
14 {
15 #ifdef CUDNN_MODE
16 template<typename T>
17 class CCuda_ComputationalGraphTransConv2D_Cudnn;
18 #endif
19
20 template <typename T>
21 class FL_EXPORT CComputationalGraphTransConv2D : public CComputationalGraph<T>
22 {
23 protected:
24 CComputationalGraphTransConv2D();
25 CComputationalGraphTransConv2D(const CComputationalGraphTransConv2D<T>& cg);
26
27 public:
28 // Operand1
29 // 4th dim : Batch size
30 // 3rd dim : Channels
31 // 2nd dim : Height
32 // 1st dim : Width
33
34 // Kernel
35 // 4th dim : Num of kernels
36 // 3rd dim : Channels
37 // 2nd dim : Height
38 // 1st dim : Width
39
40 // Bias
41 // 1st dim : Num of kernels
42 CComputationalGraphTransConv2D(const CComputationalBase<T>& cbOperand1, const CTensor<T>& tsrKernel, int64_t i64StrideX = 1, int64_t i64StrideY = 1, int64_t i64PaddingX = 0, int64_t i64PaddingY = 0, int64_t i64OutputPaddingX = 0, int64_t i64OutputPaddingY = 0, int64_t i64DilationX = 1, int64_t i64DilationY = 1, int64_t i64GroupCount = 1);
43 virtual ~CComputationalGraphTransConv2D();
44
45 virtual CTensor<T>& Forward() override;
46 virtual CTensor<T>* Backward() override;
47 virtual CComputationalBase<T>* Clone() const override;
48 virtual const CResult PrintNodeParamInfo() const override;
49
50 virtual const CResult GetBinaryData(Base::CFLData& fldBinary, bool bSuperClass = false, int32_t i32Version = -1, bool bDumpMode = false) const override;
51 virtual const CResult GetBinaryData(Base::CFLData* pFldBinary, bool bSuperClass = false, int32_t i32Version = -1, bool bDumpMode = false) const override;
52
53 virtual const CResult SetBinaryData(const Base::CFLData& fldBinary, int64_t* pI64Offset = nullptr) override;
54 virtual const CResult SetBinaryData(const Base::CFLData* pFldBinary, int64_t* pI64Offset = nullptr) override;
55
56 virtual const std::vector<int64_t>& GetEstimatedShape(bool bRecursive = true) const override;
57
58 virtual int64_t GetRequiredTemporaryMemory(bool bTraining = false, bool bRecursively = true, int64_t i64BatchSize = 1, int64_t i64MemoryIndex = 0) const override;
59 virtual int64_t GetRequiredDedicatedMemory(bool bTraining = false, bool bRecursively = true, int64_t i64BatchSize = 1) const override;
60
61
62 DeclareGetClassType();
63 SupportToDuplicateObjectWithoutCreateNewObject(CComputationalGraphTransConv2D, *this);
64
65 protected:
66 /*
67 virtual const CResult DerivativeKernel(bool bAddGradient);
68 virtual const CResult DerivativeImage(bool bAddGradient);
69 */
70 int64_t m_i64StrideX;
71 int64_t m_i64StrideY;
72 int64_t m_i64PaddingX;
73 int64_t m_i64PaddingY;
74 int64_t m_i64OutputPaddingX;
75 int64_t m_i64OutputPaddingY;
76 int64_t m_i64DilationX;
77 int64_t m_i64DilationY;
78 int64_t m_i64GroupCount;
79
80 CTensor<T> m_tsrPaddingX;
81 CTensor<T> m_tsrPaddingY;
82 CTensor<T> m_tsrInputTranspose;
83 CTensor<T> m_tsrKernelTranspose;
84 CTensor<T> m_tsrDerivativeTranspose;
85
86 CBackendConv2D<T> m_backendConv2D;
87
88 #ifdef CUDNN_MODE
89 #ifdef CUDNN_MODE
90 CCuda_ComputationalGraphTransConv2D_Cudnn<T>* m_pCudnn;
91 #endif
92 #endif
93
94 public:
95 DeclareGetSignletonObject(CComputationalGraphTransConv2D);
96 };
97
98 #define CCGFTransConv2D(...) (*(new CComputationalGraphTransConv2D<float>(__VA_ARGS__)))
99 #define CCGDTransConv2D(...) (*(new CComputationalGraphTransConv2D<double>(__VA_ARGS__)))
100
101 #define CCGTTransConv2D(T, ...) (*(new CComputationalGraphTransConv2D<T>(__VA_ARGS__)))
102 }
103}
104
105#endif
Definition AlgorithmAIBase.h:18