载入中...
搜索中...
未找到
Graph.h
浏览该文件的文档.
1#ifndef EVE_TENSOR_GRAPH_H
2#define EVE_TENSOR_GRAPH_H
3
4#include "tensor/Tensor.h"
5
6#include <cstdint>
7#include <memory>
8#include <string>
9#include <vector>
10
11namespace eve::tensor {
12
13class Tensor;
14class TF;
15class Func;
16class GpuProgram;
17struct OptimizedGraph;
18
19enum class OpType : uint8_t {
20 Placeholder = 0,
21 Const,
22 // binary (broadcast-capable)
23 Add,
24 Sub,
26 Divide,
27 // scalar / unary elementwise
32 Neg,
33 Abs,
34 Sqrt,
35 Exp,
36 Log,
37 Sin,
38 Cos,
39 Tanh,
40 Relu,
41 Sigmoid,
42 Gelu,
43 Silu,
45 Clamp,
48 Where,
49 // neural / speech / terrain ops
50 MatMul,
52 Permute,
53 Reshape,
54 Flatten,
55 Softmax,
58 RMSNorm,
59 Conv1d,
60 Conv2d,
64 Concat,
65 Slice,
70 ArgMax,
71 Cast,
74};
75
76struct GraphNode {
78 int dims[Tensor::kMaxRank] = {0, 0, 0, 0, 0, 0};
79 int rank = 0;
80 int size = 0;
81 int in0 = -1;
82 int in1 = -1;
83 int in2 = -1;
84 int in3 = -1;
85 int in4 = -1;
86 float s0 = 0.f;
87 float s1 = 0.f;
88 float s2 = 0.f;
89 float s3 = 0.f;
90 // generic int attributes: axis / stride / pad / begin / end / mode / keepdims ...
91 int i0 = 0;
92 int i1 = 0;
93 int i2 = 0;
94 int i3 = 0;
95 int perm[Tensor::kMaxRank] = {0, 1, 2, 3, 4, 5};
96 int permRank = 0;
98 int dtype = static_cast<int>(DType::Float32);
99 std::vector<float> constData;
100 // Weight-quantized const payload (dtype = Fp16/Fp8E4M3/Fp4E2M1/Int8/Int4).
101 std::vector<uint8_t> constBytes;
102 std::vector<float> constScales; // per-group scales (int8/int4)
103 int qGroup = 0; // elements per scale group
104};
105
106class Graph {
107public:
109 const GraphNode &node(int id) const { return nodes_[static_cast<size_t>(id)]; }
110 GraphNode &node(int id) { return nodes_[static_cast<size_t>(id)]; }
111 int nodeCount() const { return int(nodes_.size()); }
112 const std::vector<GraphNode> &nodes() const { return nodes_; }
113 std::vector<GraphNode> &nodes() { return nodes_; }
114
115 static int product(const int *dims, int rank);
116
117private:
118 std::vector<GraphNode> nodes_;
119};
120
125class Func {
126public:
127 explicit Func(TF *owner);
128 ~Func();
129
130 Tensor *input1(int d0);
131 Tensor *input2(int d0, int d1);
132 Tensor *input3(int d0, int d1, int d2);
133 Tensor *input4(int d0, int d1, int d2, int d3);
134 Tensor *input5(int d0, int d1, int d2, int d3, int d4);
135 Tensor *input6(int d0, int d1, int d2, int d3, int d4, int d5);
136
137 void setOutput(Tensor *t);
138
139 class CompiledFunction *compile();
140
141 Graph &graph() { return graph_; }
142 const Graph &graph() const { return graph_; }
143 TF *owner() const { return owner_; }
144 bool isTracing() const { return tracing_; }
145 int outputNode() const { return outputNode_; }
146 int placeholderCount() const { return placeholderCount_; }
147
149 int ensureNode(const Tensor *t);
150
152 Tensor *emitUnaryScalar(OpType type, const Tensor *x, float s0, float s1 = 0.f);
153 Tensor *emitBinary(OpType type, const Tensor *a, const Tensor *b);
154 Tensor *emitTernary(OpType type, const Tensor *a, const Tensor *b, const Tensor *c);
155 Tensor *emitMatMul(const Tensor *a, const Tensor *b);
156 Tensor *emitTranspose(const Tensor *x);
157 Tensor *emitPermute(const Tensor *x, const int *order, int rank);
158 Tensor *emitReshape(const Tensor *x, const int *dims, int rank);
159 Tensor *emitFill(const int *dims, int rank, float value);
160
161 Tensor *emitSoftmax(const Tensor *x, int axis, bool logMode);
162 Tensor *emitLayerNorm(const Tensor *x, const Tensor *scale, const Tensor *bias, float eps);
163 Tensor *emitRMSNorm(const Tensor *x, const Tensor *scale, float eps);
164 Tensor *emitConv1d(const Tensor *x, const Tensor *w, const Tensor *bias, int stride, int pad);
165 Tensor *emitConv2d(const Tensor *x, const Tensor *w, const Tensor *bias, int stride, int pad);
166 Tensor *emitPool(OpType type, const Tensor *x, int ksize, int stride, int pad);
167 Tensor *emitEmbedding(const Tensor *table, const Tensor *indices);
168 Tensor *emitConcat(const Tensor *const *ins, int n, int axis);
169 Tensor *emitSlice(const Tensor *x, int axis, int begin, int end);
170 Tensor *emitReduce(OpType type, const Tensor *x, int axis, bool keepDims);
171 Tensor *emitArgMax(const Tensor *x, int axis, bool keepDims);
172 Tensor *emitCast(const Tensor *x, DType dtype);
173 Tensor *emitSdpa(const Tensor *q, const Tensor *k, const Tensor *v, const Tensor *mask,
174 float scale);
175 Tensor *emitResize2d(const Tensor *x, int outH, int outW, int mode);
176
177private:
178 Tensor *makeSymbolicFromNode(int nodeId);
179 GraphNode makeShapeNode(OpType type, const int *dims, int rank);
180
181 TF *owner_ = nullptr;
182 Graph graph_;
183 int outputNode_ = -1;
184 int placeholderCount_ = 0;
185 bool tracing_ = true;
186};
187
192public:
195
196 Tensor *run0();
197 Tensor *run1(Tensor *in0);
198 Tensor *run2(Tensor *in0, Tensor *in1);
199 Tensor *run3(Tensor *in0, Tensor *in1, Tensor *in2);
200 Tensor *run4(Tensor *in0, Tensor *in1, Tensor *in2, Tensor *in3);
201 Tensor *run5(Tensor *in0, Tensor *in1, Tensor *in2, Tensor *in3, Tensor *in4);
202 Tensor *run6(Tensor *in0, Tensor *in1, Tensor *in2, Tensor *in3, Tensor *in4, Tensor *in5);
203
204 int getPlaceholderCount() const { return placeholderCount_; }
205 std::string getDevice() const { return device_; }
206
208
209private:
210 Tensor *runWithFeeds(Tensor *const *feeds, int nFeeds);
211 void executeNode(int nodeId, std::vector<std::vector<float>> &bufs) const;
212
213 Graph graph_;
214 std::vector<int> order_;
215 std::unique_ptr<OptimizedGraph> optimized_;
216 int outputNode_ = -1;
217 int placeholderCount_ = 0;
218 std::string device_ = "cpu";
220 std::unique_ptr<GpuProgram> gpuProgram_;
221};
222
223} // namespace eve::tensor
224
225#endif // EVE_TENSOR_GRAPH_H
std::string value
std::string type
std::string id
int x
Definition Grass.cpp:135
glm::vec3 n
Definition Grass.cpp:64
int w
uint32_t a
uint32_t b
uint32_t c
SettlementPipeline::Stage fn
int v
float scale
Definition TreeMesh.cpp:122
Optimized / scheduled graph ready to run with feeds.
Definition Graph.h:191
Tensor * run1(Tensor *in0)
Definition Graph.cpp:613
Tensor * run2(Tensor *in0, Tensor *in1)
Definition Graph.cpp:618
Tensor * run5(Tensor *in0, Tensor *in1, Tensor *in2, Tensor *in3, Tensor *in4)
Definition Graph.cpp:633
Tensor * run4(Tensor *in0, Tensor *in1, Tensor *in2, Tensor *in3)
Definition Graph.cpp:628
std::string getDevice() const
Definition Graph.h:205
int getPlaceholderCount() const
Definition Graph.h:204
Tensor * run3(Tensor *in0, Tensor *in1, Tensor *in2)
Definition Graph.cpp:623
Tensor * run6(Tensor *in0, Tensor *in1, Tensor *in2, Tensor *in3, Tensor *in4, Tensor *in5)
Definition Graph.cpp:638
static CompiledFunction * fromFunc(Func *fn)
Definition Graph.cpp:590
Trace builder — TF2 tf.function analogue (tf.func in scripts). While active, TF ops record into this ...
Definition Graph.h:125
const Graph & graph() const
Definition Graph.h:142
Tensor * emitReduce(OpType type, const Tensor *x, int axis, bool keepDims)
Definition Graph.cpp:470
Tensor * emitCast(const Tensor *x, DType dtype)
Definition Graph.cpp:523
Tensor * emitPermute(const Tensor *x, const int *order, int rank)
Definition Graph.cpp:241
class CompiledFunction * compile()
Definition Graph.cpp:583
Tensor * input6(int d0, int d1, int d2, int d3, int d4, int d5)
Definition Graph.cpp:118
Tensor * emitResize2d(const Tensor *x, int outH, int outW, int mode)
Definition Graph.cpp:563
Tensor * input4(int d0, int d1, int d2, int d3)
Definition Graph.cpp:106
Tensor * emitConcat(const Tensor *const *ins, int n, int axis)
Definition Graph.cpp:418
Tensor * input1(int d0)
Definition Graph.cpp:88
int outputNode() const
Definition Graph.h:145
TF * owner() const
Definition Graph.h:143
Tensor * input5(int d0, int d1, int d2, int d3, int d4)
Definition Graph.cpp:112
Tensor * emitPool(OpType type, const Tensor *x, int ksize, int stride, int pad)
Definition Graph.cpp:384
Tensor * emitSdpa(const Tensor *q, const Tensor *k, const Tensor *v, const Tensor *mask, float scale)
Definition Graph.cpp:533
Tensor * emitBinary(OpType type, const Tensor *a, const Tensor *b)
Definition Graph.cpp:177
Tensor * emitSoftmax(const Tensor *x, int axis, bool logMode)
Definition Graph.cpp:271
Tensor * emitConv1d(const Tensor *x, const Tensor *w, const Tensor *bias, int stride, int pad)
Definition Graph.cpp:331
Tensor * emitTernary(OpType type, const Tensor *a, const Tensor *b, const Tensor *c)
Definition Graph.cpp:193
Tensor * emitFill(const int *dims, int rank, float value)
Definition Graph.cpp:148
void setOutput(Tensor *t)
Definition Graph.cpp:124
Graph & graph()
Definition Graph.h:141
Tensor * emitEmbedding(const Tensor *table, const Tensor *indices)
Definition Graph.cpp:400
Tensor * emitConv2d(const Tensor *x, const Tensor *w, const Tensor *bias, int stride, int pad)
Definition Graph.cpp:357
Tensor * emitArgMax(const Tensor *x, int axis, bool keepDims)
Definition Graph.cpp:496
Tensor * emitMatMul(const Tensor *a, const Tensor *b)
Definition Graph.cpp:206
Tensor * emitUnaryScalar(OpType type, const Tensor *x, float s0, float s1=0.f)
Definition Graph.cpp:165
Tensor * input2(int d0, int d1)
Definition Graph.cpp:94
Tensor * emitReshape(const Tensor *x, const int *dims, int rank)
Definition Graph.cpp:259
Tensor * emitSlice(const Tensor *x, int axis, int begin, int end)
Definition Graph.cpp:452
int placeholderCount() const
Definition Graph.h:146
Tensor * emitUnary(OpType type, const Tensor *x)
Definition Graph.cpp:155
Tensor * emitTranspose(const Tensor *x)
Definition Graph.cpp:233
bool isTracing() const
Definition Graph.h:144
Tensor * input3(int d0, int d1, int d2)
Definition Graph.cpp:100
int ensureNode(const Tensor *t)
Ensure tensor is a node in this graph (Const-capture if eager).
Definition Graph.cpp:129
Tensor * emitLayerNorm(const Tensor *x, const Tensor *scale, const Tensor *bias, float eps)
Definition Graph.cpp:292
Tensor * emitRMSNorm(const Tensor *x, const Tensor *scale, float eps)
Definition Graph.cpp:314
GraphNode & node(int id)
Definition Graph.h:110
int addNode(GraphNode node)
Definition Graph.cpp:25
const GraphNode & node(int id) const
Definition Graph.h:109
const std::vector< GraphNode > & nodes() const
Definition Graph.h:112
std::vector< GraphNode > & nodes()
Definition Graph.h:113
int nodeCount() const
Definition Graph.h:111
static int product(const int *dims, int rank)
Definition Graph.cpp:16
TF2-like namespace module. Script: tf <- eve.TF(); Default eager; tf.func() traces a graph for compil...
Definition TF.h:18
float32 / int32 tensor (rank 1–6), row-major. Eager: owns a buffer. Symbolic: node in a Func graph (n...
Definition Tensor.h:43
static constexpr int kMaxRank
Definition Tensor.h:45
DType
Tensor element types.
Definition Tensor.h:22
std::vector< float > constData
Definition Graph.h:99
int perm[Tensor::kMaxRank]
Definition Graph.h:95
int dims[Tensor::kMaxRank]
Definition Graph.h:78
std::vector< float > constScales
Definition Graph.h:102
std::vector< uint8_t > constBytes
Definition Graph.h:101