3#if _MSC_VER >= 1900 && defined(_M_X64)
5#include "DefinitionsAI.h"
6#include "ComputationalGraphUtilities.h"
16 class CComputationalGraph;
24 virtual ~CComputationalBase();
26 virtual bool IsInitialized()
const;
27 virtual const CResult SetInitialized(
bool bSet);
29 virtual const CResult
Load(
const wchar_t* pWcsFileName);
30 virtual const CResult
Save(
const wchar_t* pWcsFileName);
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;
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);
38 virtual const CResult
Clear();
40 const int64_t GetObjectNumber()
const;
42 CComputationalBase<T>& SetID(
const wchar_t* pWcsName);
43 CComputationalBase<T>& ID(
const wchar_t* pWcsName);
44 const wchar_t* GetID()
const;
46 virtual void Throw(
const CResult& res,
const wchar_t* pWcsExtraMessage =
nullptr)
const override;
48 CComputationalBase<T>& EnableReference(
bool bReference);
49 bool IsReference()
const;
51 const CResult SetValueAttribute(EValueAttribute eType);
52 EValueAttribute GetValueAttribute()
const;
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;
65 virtual CTensor<T>& Evaluate() = 0;
66 virtual CTensor<T>& Forward() = 0;
67 virtual CTensor<T>* Backward() = 0;
69 virtual CTensor<T>& GetValue()
const;
70 virtual CTensor<T>* GetDerivative()
const;
71 virtual void ClearDerivativesRecursive();
72 virtual void ResetDerivativesRecursive();
74 virtual CComputationalBase<T>* Clone()
const = 0;
76 virtual const CResult Swap(CComputationalBase<T>& cbSwap);
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;
81 virtual CComputationalBase<T>* GetOperand(int64_t i64Index)
const;
82 virtual bool IsIntrinsicOperand(int64_t i64Index)
const;
83 virtual int64_t GetOperandCount()
const;
85 virtual int64_t GetGeneration()
const;
87 virtual bool IsTrainingModeEnabled()
const;
88 virtual const CResult EnableTrainingMode(
bool bMode,
bool bRecursively =
true);
90 virtual const CResult PrintGraphInfo(
bool bRecursively =
true,
bool bShapeOrderAsc =
false,
bool bIncludeTensors =
false)
const;
91 virtual const CResult PrintNodeParamInfo()
const;
93 virtual bool IsTensorCoreAvailable()
const;
95 virtual int64_t GetAddGradientCount()
const;
96 virtual const CResult SetAddGradientCount(int64_t i64AddGradientCount);
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;
101 virtual const CResult EnableRetainedValue(
bool bRetained);
102 virtual bool IsRetainedEnabled()
const;
104 virtual bool IsInplace()
const;
106 virtual int64_t GetParentNodeCount()
const;
107 virtual CComputationalBase<T>* GetParentNode(int64_t i64Index)
const;
109 virtual const CResult SetDeviceIndex(int32_t i32DeviceIndex);
110 virtual int32_t GetDeviceIndex()
const;
112 virtual const CResult CreateValue();
113 virtual const CResult ClearValue();
115 virtual const CResult CreateDerivative();
116 virtual const CResult ClearDerivative();
118 const CResult SetValuePtr(CTensor<T>* pTsrValue);
119 const CResult SetDerivativePtr(CTensor<T>* pTsrDerivative);
121 virtual CTensor<T>* GetValuePtr();
123 static CComputationalBase<T>* GetSingletonObject();
125 virtual const CResult SetProcessingUnit(
const Base::CProcessingUnitBase& puBase)
override;
126 virtual const CResult SetProcessingUnit(
const Base::CProcessingUnitBase* pPuBase)
override;
128 virtual const CResult EnableTensorCore(
bool bTensorCore =
true);
129 virtual bool IsTensorCoreEnabled()
const;
131 SupportToDuplicateAbstractObject(CComputationalBase);
133 const CResult InternalAssign(
const CComputationalBase<T>& cb);
135 virtual const CResult BeginForwardPerformanceCheck();
136 virtual const CResult EndForwardPerformanceCheck();
138 virtual const CResult BeginBackwardPerformanceCheck();
139 virtual const CResult EndBackwardPerformanceCheck();
143 ENodeType m_eNodeType;
144 EDataType m_eDataType;
145 EValueAttribute m_eValueAttribute;
146 ENodeOperator m_eNodeOperator;
148 bool m_bTrainingMode;
149 int64_t m_i64ObjectNumber;
152 bool m_bTensorCoreEnabled;
153 int64_t m_i64AddGradientCount;
154 int64_t m_i64Generation;
155 int32_t m_i32DeviceIndex;
157 CTensor<T>* m_pTsrValue;
158 CTensor<T>* m_pTsrDerivative;
160 Base::CPerformanceCounter m_perfForward;
161 Base::CPerformanceCounter m_perfBackward;
163 std::vector<int64_t>& m_vctShape;
164 std::vector<int64_t>& m_vctShapeAsc;
166 std::vector<CComputationalBase<T>*>& m_vctParentNodes;
168 int64_t m_i64CommonVersion;
169 bool m_bNonTrainingMemoryInheritance;
171 friend class CComputationalGraph<T>;
172 friend class CComputationalGraphUtilities<T>;
175 typedef CComputationalBase<float> CComputationalBaseF;
176 typedef CComputationalBase<double> CComputationalBaseD;
178 typedef CComputationalBase<float> CCBF;
179 typedef CComputationalBase<double> CCBD;
181 template <
typename T>
182 using CCB = CComputationalBase<T>;
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