载入中...
搜索中...
未找到
Graph.h
浏览该文件的文档.
1
2#include "common/Export.h"
3#ifndef EVE_TENSOR_GRAPH_H
4#define EVE_TENSOR_GRAPH_H
5
6#include "tensor/Tensor.h"
7
8#include <cstdint>
9#include <memory>
10#include <string>
11#include <vector>
12
13namespace eve::tensor {
14
15class Tensor;
16class TF;
17class Func;
18class GpuProgram;
19struct OptimizedGraph;
20
22enum class OpType : uint8_t {
23 Placeholder = 0,
24 Const,
25 // binary (broadcast-capable)
26 Add,
27 Sub,
29 Divide,
30 // scalar / unary elementwise
35 Neg,
36 Abs,
37 Sqrt,
38 Exp,
39 Log,
40 Sin,
41 Cos,
42 Tanh,
43 Relu,
44 Sigmoid,
45 Gelu,
46 Silu,
48 Clamp,
51 Where,
52 // neural / speech / terrain ops
53 MatMul,
55 Permute,
56 Reshape,
57 Flatten,
58 Softmax,
61 RMSNorm,
62 Conv1d,
63 Conv2d,
67 Concat,
68 Slice,
73 ArgMax,
74 Cast,
77};
78
80struct GraphNode {
82 int dims[Tensor::kMaxRank] = {0, 0, 0, 0, 0, 0};
83 int rank = 0;
84 int size = 0;
85 int in0 = -1;
86 int in1 = -1;
87 int in2 = -1;
88 int in3 = -1;
89 int in4 = -1;
90 float s0 = 0.f;
91 float s1 = 0.f;
92 float s2 = 0.f;
93 float s3 = 0.f;
94 // generic int attributes: axis / stride / pad / begin / end / mode / keepdims ...
95 int i0 = 0;
96 int i1 = 0;
97 int i2 = 0;
98 int i3 = 0;
99 int perm[Tensor::kMaxRank] = {0, 1, 2, 3, 4, 5};
100 int permRank = 0;
102 int dtype = static_cast<int>(DType::Float32);
103 std::vector<float> constData;
104 // Weight-quantized const payload (dtype = Fp16/Fp8E4M3/Fp4E2M1/Int8/Int4).
105 std::vector<uint8_t> constBytes;
106 std::vector<float> constScales; // per-group scales (int8/int4)
107 int qGroup = 0; // elements per scale group
108};
109
112public:
114 int addNode(GraphNode node);
116 const GraphNode &node(int id) const { return nodes_[static_cast<size_t>(id)]; }
118 GraphNode &node(int id) { return nodes_[static_cast<size_t>(id)]; }
120 int nodeCount() const { return int(nodes_.size()); }
122 const std::vector<GraphNode> &nodes() const { return nodes_; }
124 std::vector<GraphNode> &nodes() { return nodes_; }
125
127 static int product(const int *dims, int rank);
128
129private:
130 std::vector<GraphNode> nodes_;
131};
132
138public:
140 explicit Func(TF *owner);
142 ~Func();
143
145 Tensor *input1(int d0);
147 Tensor *input2(int d0, int d1);
149 Tensor *input3(int d0, int d1, int d2);
151 Tensor *input4(int d0, int d1, int d2, int d3);
153 Tensor *input5(int d0, int d1, int d2, int d3, int d4);
155 Tensor *input6(int d0, int d1, int d2, int d3, int d4, int d5);
156
158 void setOutput(Tensor *t);
159
161 class CompiledFunction *compile();
162
164 Graph &graph() { return graph_; }
166 const Graph &graph() const { return graph_; }
168 TF *owner() const { return owner_; }
170 bool isTracing() const { return tracing_; }
172 int outputNode() const { return outputNode_; }
174 int placeholderCount() const { return placeholderCount_; }
175
177 int ensureNode(const Tensor *t);
178
180 Tensor *emitUnary(OpType type, const Tensor *x);
182 Tensor *emitUnaryScalar(OpType type, const Tensor *x, float s0, float s1 = 0.f);
184 Tensor *emitBinary(OpType type, const Tensor *a, const Tensor *b);
186 Tensor *emitTernary(OpType type, const Tensor *a, const Tensor *b, const Tensor *c);
188 Tensor *emitMatMul(const Tensor *a, const Tensor *b);
190 Tensor *emitTranspose(const Tensor *x);
192 Tensor *emitPermute(const Tensor *x, const int *order, int rank);
194 Tensor *emitReshape(const Tensor *x, const int *dims, int rank);
196 Tensor *emitFill(const int *dims, int rank, float value);
197
199 Tensor *emitSoftmax(const Tensor *x, int axis, bool logMode);
201 Tensor *emitLayerNorm(const Tensor *x, const Tensor *scale, const Tensor *bias, float eps);
203 Tensor *emitRMSNorm(const Tensor *x, const Tensor *scale, float eps);
205 Tensor *emitConv1d(const Tensor *x, const Tensor *w, const Tensor *bias, int stride, int pad);
207 Tensor *emitConv2d(const Tensor *x, const Tensor *w, const Tensor *bias, int stride, int pad);
209 Tensor *emitPool(OpType type, const Tensor *x, int ksize, int stride, int pad);
211 Tensor *emitEmbedding(const Tensor *table, const Tensor *indices);
213 Tensor *emitConcat(const Tensor *const *ins, int n, int axis);
215 Tensor *emitSlice(const Tensor *x, int axis, int begin, int end);
217 Tensor *emitReduce(OpType type, const Tensor *x, int axis, bool keepDims);
219 Tensor *emitArgMax(const Tensor *x, int axis, bool keepDims);
221 Tensor *emitCast(const Tensor *x, DType dtype);
223 Tensor *emitSdpa(const Tensor *q, const Tensor *k, const Tensor *v, const Tensor *mask,
224 float scale);
226 Tensor *emitResize2d(const Tensor *x, int outH, int outW, int mode);
227
228private:
229 Tensor *makeSymbolicFromNode(int nodeId);
230 GraphNode makeShapeNode(OpType type, const int *dims, int rank);
231
232 TF *owner_ = nullptr;
233 Graph graph_;
234 int outputNode_ = -1;
235 int placeholderCount_ = 0;
236 bool tracing_ = true;
237};
238
243public:
248
250 Tensor *run0();
252 Tensor *run1(Tensor *in0);
254 Tensor *run2(Tensor *in0, Tensor *in1);
256 Tensor *run3(Tensor *in0, Tensor *in1, Tensor *in2);
258 Tensor *run4(Tensor *in0, Tensor *in1, Tensor *in2, Tensor *in3);
260 Tensor *run5(Tensor *in0, Tensor *in1, Tensor *in2, Tensor *in3, Tensor *in4);
262 Tensor *run6(Tensor *in0, Tensor *in1, Tensor *in2, Tensor *in3, Tensor *in4, Tensor *in5);
263
265 int getPlaceholderCount() const { return placeholderCount_; }
267 std::string getDevice() const { return device_; }
268
270 static CompiledFunction *fromFunc(Func *fn);
271
272private:
273 Tensor *runWithFeeds(Tensor *const *feeds, int nFeeds);
274 void executeNode(int nodeId, std::vector<std::vector<float>> &bufs) const;
275
276 Graph graph_;
277 std::vector<int> order_;
278 std::unique_ptr<OptimizedGraph> optimized_;
279 int outputNode_ = -1;
280 int placeholderCount_ = 0;
281 std::string device_ = "cpu";
283 std::unique_ptr<GpuProgram> gpuProgram_;
284};
285
286} // namespace eve::tensor
287
288#endif // EVE_TENSOR_GRAPH_H
double value
float w
Definition AnimClip.cpp:738
float x
Definition AnimClip.cpp:738
int mask
std::string nodeId
#define EVENGINE_API_DOMAINS
Definition Export.h:110
glm::vec3 n
Definition Grass.cpp:63
std::array< double, 10 > q
std::vector< std::uint32_t > indices
float v
std::int32_t c
std::array< float, 3 > scale
MeleePoint3 b
Definition MeleeHit.cpp:41
MeleePoint3 a
Definition MeleeHit.cpp:40
std::vector< std::int32_t > order
std::string id
Definition PlayHost.cpp:108
float begin
float t
const RoadNode * node
float bias
Optimized / scheduled graph ready to run with feeds.
Definition Graph.h:242
std::string getDevice() const
Returns the device.
Definition Graph.h:267
CompiledFunction()
Compiled function.
int getPlaceholderCount() const
Returns the placeholder count.
Definition Graph.h:265
~CompiledFunction()
Compiled function.
Trace builder — TF2 tf.function analogue (tf.func in scripts). While active, TF ops record into this ...
Definition Graph.h:137
const Graph & graph() const
Graph.
Definition Graph.h:166
int outputNode() const
Output node.
Definition Graph.h:172
TF * owner() const
Owner.
Definition Graph.h:168
Graph & graph()
Graph.
Definition Graph.h:164
int placeholderCount() const
Placeholder count.
Definition Graph.h:174
bool isTracing() const
True when tracing.
Definition Graph.h:170
EVENGINE_API_DOMAINS public API.
Definition Graph.h:111
GraphNode & node(int id)
Node.
Definition Graph.h:118
const GraphNode & node(int id) const
Node.
Definition Graph.h:116
const std::vector< GraphNode > & nodes() const
Nodes.
Definition Graph.h:122
std::vector< GraphNode > & nodes()
Nodes.
Definition Graph.h:124
int nodeCount() const
Node count.
Definition Graph.h:120
TF2-like namespace module. Script: tf <- eve.TF(); Default eager; tf.func() traces a graph for compil...
Definition TF.h:20
float32 / int32 tensor (rank 1–6), row-major. Eager: owns a buffer. Symbolic: node in a Func graph (n...
Definition Tensor.h:48
static constexpr int kMaxRank
Definition Tensor.h:50
DType
Tensor element types.
Definition Tensor.h:24
OpType
OpType public API.
Definition Graph.h:22
SettlementPipeline::Stage fn
GraphNode public API.
Definition Graph.h:80
std::vector< float > constData
Definition Graph.h:103
int perm[Tensor::kMaxRank]
Definition Graph.h:99
int dims[Tensor::kMaxRank]
Definition Graph.h:82
std::vector< float > constScales
Definition Graph.h:106
std::vector< uint8_t > constBytes
Definition Graph.h:105
uint32_t pad[2]