FLImaging 7.8.25.3
ComputationalGraphLSTM.h
1#pragma once
2
3#if _MSC_VER >= 1900 && defined(_M_X64)
4
5#include "ComputationalGraph.h"
6#include "ComputationalGraphRNN.h"
7#include "Tensor.h"
8#include "BackendRNN.h"
9
10#include <vector>
11
12namespace FLImaging
13{
14 namespace AI
15 {
16 template <typename T>
17 class FL_EXPORT CComputationalGraphLSTM : public CComputationalGraphRNN<T>
18 {
19 protected:
20 CComputationalGraphLSTM();
21 CComputationalGraphLSTM(const CComputationalGraphLSTM<T>& cg);
22
23 public:
24
25 CComputationalGraphLSTM(const CComputationalBase<T>& cbOperand, int64_t i64HiddenSize, bool bBatchFirst = false, bool bBias = true, bool bBidirectional = false);
26 virtual ~CComputationalGraphLSTM();
27
28 virtual CTensor<T>& Forward() override;
29 virtual CTensor<T>* Backward() override;
30 virtual CComputationalBase<T>* Clone() const override;
31
32 virtual const CResult GetBinaryData(Base::CFLData& fldBinary, bool bSuperClass = false, int32_t i32Version = -1, bool bDumpMode = false) const override;
33 virtual const CResult GetBinaryData(Base::CFLData* pFldBinary, bool bSuperClass = false, int32_t i32Version = -1, bool bDumpMode = false) const override;
34 virtual const CResult SetBinaryData(const Base::CFLData& fldBinary, int64_t* pI64Offset = nullptr) override;
35 virtual const CResult SetBinaryData(const Base::CFLData* pFldBinary, int64_t* pI64Offset = nullptr) override;
36
37 virtual int64_t GetRequiredDedicatedMemory(bool bTraining = false, bool bRecursively = true, int64_t i64BatchSize = 1) const override;
38 virtual int64_t GetRequiredTemporaryMemory(bool bTraining = false, bool bRecursively = true, int64_t i64BatchSize = 1, int64_t i64MemoryIndex = 0) const override;
39
40 DeclareGetClassType();
41 SupportToDuplicateObjectWithoutCreateNewObject(CComputationalGraphLSTM, *this);
42
43 protected:
44
45 CTensor<T>* m_pTsrAllGateResult;
46 CTensor<T>* m_pTsrReverseAllGateResult;
47 CTensor<T>* m_pTsrCell;
48 CTensor<T>* m_pTsrReverseCell;
49
50 public:
51 DeclareGetSignletonObject(CComputationalGraphLSTM);
52 };
53
54 #define CCGFLSTM(...) (*(new CComputationalGraphLSTM<float>(__VA_ARGS__)))
55 #define CCGDLSTM(...) (*(new CComputationalGraphLSTM<double>(__VA_ARGS__)))
56
57 #define CCGTLSTM(T, ...) (*(new CComputationalGraphLSTM<T>(__VA_ARGS__)))
58 }
59}
60
61#endif
Definition AlgorithmAIBase.h:18