FLImaging 7.8.25.3
ComputationalBase.h
1#pragma once
2
3#if _MSC_VER >= 1900 && defined(_M_X64)
4
5#include "DefinitionsAI.h"
6#include "ComputationalGraphUtilities.h"
7
8namespace FLImaging
9{
10 namespace AI
11 {
12 template <typename T>
13 class CTensor;
14
15 template <typename T>
16 class CComputationalGraph;
17
18 template <typename T>
19 class FL_EXPORT CComputationalBase : public CAlgorithmAIBase
20 {
21
22 public:
23 CComputationalBase();
24 virtual ~CComputationalBase();
25
26 virtual bool IsInitialized() const;
27 virtual const CResult SetInitialized(bool bSet);
28
29 virtual const CResult Load(const wchar_t* pWcsFileName);
30 virtual const CResult Save(const wchar_t* pWcsFileName);
31
32 virtual const CResult GetBinaryData(Base::CFLData& fldBinary, bool bSuperClass = false, int32_t i32Version = -1, bool bDumpMode = false) const;
33 virtual const CResult GetBinaryData(Base::CFLData* pFldBinary, bool bSuperClass = false, int32_t i32Version = -1, bool bDumpMode = false) const;
34
35 virtual const CResult SetBinaryData(const Base::CFLData& fldBinary, int64_t* pI64Offset = nullptr);
36 virtual const CResult SetBinaryData(const Base::CFLData* pFldBinary, int64_t* pI64Offset = nullptr);
37
38 virtual const CResult Clear();
39
40 const int64_t GetObjectNumber() const;
41
42 CComputationalBase<T>& SetID(const wchar_t* pWcsName);
43 CComputationalBase<T>& ID(const wchar_t* pWcsName);
44 const wchar_t* GetID() const;
45
46 virtual void Throw(const CResult& res, const wchar_t* pWcsExtraMessage = nullptr) const override;
47
48 CComputationalBase<T>& EnableReference(bool bReference);
49 bool IsReference() const;
50
51 const CResult SetValueAttribute(EValueAttribute eType);
52 EValueAttribute GetValueAttribute() const;
53
54 ENodeType GetNodeType() const;
55 EDataType GetDataType() const;
56 ENodeOperator GetNodeOperator() const;
57 virtual const std::vector<int64_t>& GetEstimatedShape(bool bRecursive = true) const;
58 virtual CResult EnableNonTrainingMemoryInheritance(bool bEnable);
59 virtual bool IsNonTrainingMemoryInheritanceEnabled() const;
60 virtual const std::vector<int64_t>& GetShape() const;
61 virtual const std::vector<int64_t>& GetShapeAsc() const;
62 virtual const CResult ClearShape(bool bRecursive = false);
63 virtual int64_t GetNextBatchSize(int64_t i64BatchSize) const;
64
65 virtual CTensor<T>& Evaluate() = 0;
66 virtual CTensor<T>& Forward() = 0;
67 virtual CTensor<T>* Backward() = 0;
68
69 virtual CTensor<T>& GetValue() const;
70 virtual CTensor<T>* GetDerivative() const;
71 virtual void ClearDerivativesRecursive();
72 virtual void ResetDerivativesRecursive();
73
74 virtual CComputationalBase<T>* Clone() const = 0;
75
76 virtual const CResult Swap(CComputationalBase<T>& cbSwap);
77
78 virtual const CComputationalBase<T>* GetAt(const wchar_t* pWcsName) const = 0;
79 virtual const CResult FindByValueAttribute(EValueAttribute eValueAttribute, std::vector<const CComputationalBase<T>*>& vctResult) const = 0;
80
81 virtual CComputationalBase<T>* GetOperand(int64_t i64Index) const;
82 virtual bool IsIntrinsicOperand(int64_t i64Index) const;
83 virtual int64_t GetOperandCount() const;
84
85 virtual int64_t GetGeneration() const;
86
87 virtual bool IsTrainingModeEnabled() const;
88 virtual const CResult EnableTrainingMode(bool bMode, bool bRecursively = true);
89
90 virtual const CResult PrintGraphInfo(bool bRecursively = true, bool bShapeOrderAsc = false, bool bIncludeTensors = false) const;
91 virtual const CResult PrintNodeParamInfo() const;
92
93 virtual bool IsTensorCoreAvailable() const;
94
95 virtual int64_t GetAddGradientCount() const;
96 virtual const CResult SetAddGradientCount(int64_t i64AddGradientCount);
97
98 virtual int64_t GetRequiredDedicatedMemory(bool bTraining = false, bool bRecursively = true, int64_t i64BatchSize = 1) const = 0;
99 virtual int64_t GetRequiredTemporaryMemory(bool bTraining = false, bool bRecursively = true, int64_t i64BatchSize = 1, int64_t i64MemoryIndex = 0) const = 0;
100
101 virtual const CResult EnableRetainedValue(bool bRetained);
102 virtual bool IsRetainedEnabled() const;
103
104 virtual bool IsInplace() const;
105
106 virtual int64_t GetParentNodeCount() const;
107 virtual CComputationalBase<T>* GetParentNode(int64_t i64Index) const;
108
109 virtual const CResult SetDeviceIndex(int32_t i32DeviceIndex);
110 virtual int32_t GetDeviceIndex() const;
111
112 virtual const CResult CreateValue();
113 virtual const CResult ClearValue();
114
115 virtual const CResult CreateDerivative();
116 virtual const CResult ClearDerivative();
117
118 const CResult SetValuePtr(CTensor<T>* pTsrValue);
119 const CResult SetDerivativePtr(CTensor<T>* pTsrDerivative);
120
121 virtual CTensor<T>* GetValuePtr();
122
123 static CComputationalBase<T>* GetSingletonObject();
124
125 virtual const CResult SetProcessingUnit(const Base::CProcessingUnitBase& puBase) override;
126 virtual const CResult SetProcessingUnit(const Base::CProcessingUnitBase* pPuBase) override;
127
128 virtual const CResult EnableTensorCore(bool bTensorCore = true);
129 virtual bool IsTensorCoreEnabled() const;
130
131 SupportToDuplicateAbstractObject(CComputationalBase);
132 protected:
133 const CResult InternalAssign(const CComputationalBase<T>& cb);
134
135 virtual const CResult BeginForwardPerformanceCheck();
136 virtual const CResult EndForwardPerformanceCheck();
137
138 virtual const CResult BeginBackwardPerformanceCheck();
139 virtual const CResult EndBackwardPerformanceCheck();
140
141 bool m_bInitialized;
142 bool m_bReference;
143 ENodeType m_eNodeType;
144 EDataType m_eDataType;
145 EValueAttribute m_eValueAttribute;
146 ENodeOperator m_eNodeOperator;
147 wchar_t* m_pWcsName;
148 bool m_bTrainingMode;
149 int64_t m_i64ObjectNumber;
150 bool m_bRetained;
151 bool m_bInplace;
152 bool m_bTensorCoreEnabled;
153 int64_t m_i64AddGradientCount;
154 int64_t m_i64Generation;
155 int32_t m_i32DeviceIndex;
156
157 CTensor<T>* m_pTsrValue;
158 CTensor<T>* m_pTsrDerivative;
159
160 Base::CPerformanceCounter m_perfForward;
161 Base::CPerformanceCounter m_perfBackward;
162
163 std::vector<int64_t>& m_vctShape;
164 std::vector<int64_t>& m_vctShapeAsc;
165
166 std::vector<CComputationalBase<T>*>& m_vctParentNodes;
167
168 int64_t m_i64CommonVersion;
169 bool m_bNonTrainingMemoryInheritance;
170 private:
171 friend class CComputationalGraph<T>;
172 friend class CComputationalGraphUtilities<T>;
173 };
174
175 typedef CComputationalBase<float> CComputationalBaseF;
176 typedef CComputationalBase<double> CComputationalBaseD;
177
178 typedef CComputationalBase<float> CCBF;
179 typedef CComputationalBase<double> CCBD;
180
181 template <typename T>
182 using CCB = CComputationalBase<T>;
183 }
184}
185
186#endif
Processing unit AI class required by algorithm.
Definition AlgorithmAIBase.h:27
Definition AlgorithmAIBase.h:18
@ Save
Save file.
Definition DefinitionsGUI.h:303
@ Clear
Clear all the figure objects.
Definition DefinitionsGUI.h:5083
@ Load
Default Load.
Definition DefinitionsGUI.h:50