载入中...
搜索中...
未找到
Optimizer.h
浏览该文件的文档.
1
2#include "common/Export.h"
3#ifndef EVE_TENSOR_OPTIMIZER_H
4#define EVE_TENSOR_OPTIMIZER_H
5
6#include "tensor/Graph.h"
7
8#include <vector>
9
10namespace eve::tensor {
11
24enum class GroupKind : uint8_t {
26 MatMul,
27 Conv1d,
28 Conv2d,
31 Softmax,
33 RMSNorm,
34 Reduce,
35 ArgMax,
37 Concat,
38 Slice,
39 Permute,
40 Sdpa,
42 Alias,
43};
44
46struct FusedGroup {
51 std::vector<int> nodes;
53 std::vector<int> inputs;
54 int outputNode = -1;
55
56 // attributes copied from the root node for convenience
57 float s0 = 0.f, s1 = 0.f, s2 = 0.f, s3 = 0.f;
58 int i0 = 0, i1 = 0, i2 = 0, i3 = 0;
59 int perm[Tensor::kMaxRank] = {0, 1, 2, 3, 4, 5};
60 int permRank = 0;
61 int dtype = static_cast<int>(DType::Float32);
62
63 // MatMul / Conv epilogue: bias node id (or -1) + elementwise chain node ids
64 int biasNode = -1;
65 std::vector<int> epilogue;
66 bool hasScale = false;
67 bool hasBias = false;
68 bool masked = false;
69 bool logMode = false;
70};
71
79 std::vector<int> order;
81 std::vector<FusedGroup> groups;
83 std::vector<int> groupOrder;
85 std::vector<int> nodeSlot;
87 std::vector<int> slotSize;
89 std::vector<int> persistentSlots;
90 int outputNode = -1;
91};
92
103
106int groupKernelCount(const OptimizedGraph &opt);
107
108} // namespace eve::tensor
109
110#endif // EVE_TENSOR_OPTIMIZER_H
#define EVENGINE_API_DOMAINS
Definition Export.h:110
std::map< std::string, std::vector< std::string > > graph
Definition Package.cpp:59
EVENGINE_API_DOMAINS public API.
Definition Graph.h:111
static constexpr int kMaxRank
Definition Tensor.h:50
int groupKernelCount(const OptimizedGraph &opt)
Group kernel count.
OpType
OpType public API.
Definition Graph.h:22
OptimizedGraph optimizeGraph(const Graph &graph, int outputNode)
Optimize graph.
GroupKind
GroupKind public API.
Definition Optimizer.h:24
FusedGroup public API.
Definition Optimizer.h:46
std::vector< int > epilogue
Definition Optimizer.h:65
int perm[Tensor::kMaxRank]
Definition Optimizer.h:59
std::vector< int > nodes
Definition Optimizer.h:51
std::vector< int > inputs
Definition Optimizer.h:53
OptimizedGraph public API.
Definition Optimizer.h:77
std::vector< int > order
Definition Optimizer.h:79
std::vector< int > slotSize
Definition Optimizer.h:87
std::vector< int > persistentSlots
Definition Optimizer.h:89
std::vector< int > nodeSlot
Definition Optimizer.h:85
std::vector< int > groupOrder
Definition Optimizer.h:83
std::vector< FusedGroup > groups
Definition Optimizer.h:81