载入中...
搜索中...
未找到
TF.h
浏览该文件的文档.
1#pragma once
2#include "common/Export.h"
3
4
5#include "common/Module.h"
6
7#include <cstdint>
8#include <vector>
9
10namespace eve::tensor {
11
12class Tensor;
13class Func;
14class CompiledFunction;
15
21public:
24 TF();
26 ~TF() override = default;
27
29 Func *func();
30
32 void pushTrace(Func *f);
34 void popTrace(Func *f);
36 Func *tracing() const;
37
38 // --- factories (eager, or Const nodes while tracing) ---
40 Tensor *zeros1(int d0);
42 Tensor *zeros2(int d0, int d1);
44 Tensor *zeros3(int d0, int d1, int d2);
46 Tensor *zeros4(int d0, int d1, int d2, int d3);
48 Tensor *zeros5(int d0, int d1, int d2, int d3, int d4);
50 Tensor *zeros6(int d0, int d1, int d2, int d3, int d4, int d5);
51
53 Tensor *ones1(int d0);
55 Tensor *ones2(int d0, int d1);
57 Tensor *ones3(int d0, int d1, int d2);
59 Tensor *ones4(int d0, int d1, int d2, int d3);
61 Tensor *ones5(int d0, int d1, int d2, int d3, int d4);
63 Tensor *ones6(int d0, int d1, int d2, int d3, int d4, int d5);
64
66 Tensor *fill1(int d0, float value);
68 Tensor *fill2(int d0, int d1, float value);
70 Tensor *fill3(int d0, int d1, int d2, float value);
72 Tensor *fill4(int d0, int d1, int d2, int d3, float value);
73
75 Tensor *constantScalar(float value);
77 Tensor *arange(int n);
79 Tensor *linspace(float start, float end, int n);
81 Tensor *eye(int n);
82
84 Tensor *randomUniform1(int d0);
86 Tensor *randomUniform2(int d0, int d1);
88 Tensor *randomUniform3(int d0, int d1, int d2);
90 Tensor *randomUniform4(int d0, int d1, int d2, int d3);
92 Tensor *randomNormal1(int d0);
94 Tensor *randomNormal2(int d0, int d1);
96 Tensor *randomNormal3(int d0, int d1, int d2);
98 Tensor *randomNormal4(int d0, int d1, int d2, int d3);
99
100 // aliases
102 Tensor *rand1(int d0) { return randomUniform1(d0); }
104 Tensor *rand2(int d0, int d1) { return randomUniform2(d0, d1); }
106 Tensor *rand3(int d0, int d1, int d2) { return randomUniform3(d0, d1, d2); }
108 Tensor *rand4(int d0, int d1, int d2, int d3) {
110 return randomUniform4(d0, d1, d2, d3);
111 }
113 Tensor *randn1(int d0) { return randomNormal1(d0); }
115 Tensor *randn2(int d0, int d1) { return randomNormal2(d0, d1); }
117 Tensor *randn3(int d0, int d1, int d2) { return randomNormal3(d0, d1, d2); }
119 Tensor *randn4(int d0, int d1, int d2, int d3) {
121 return randomNormal4(d0, d1, d2, d3);
122 }
123
125 void setRandomSeed(uint32_t seed);
127 uint32_t getRandomSeed() const;
128
129 // --- module-level ops (TF style) ---
139 Tensor *addScalar(Tensor *a, float s);
141 Tensor *subScalar(Tensor *a, float s);
143 Tensor *mulScalar(Tensor *a, float s);
145 Tensor *divScalar(Tensor *a, float s);
171 Tensor *powScalar(Tensor *a, float exp);
173 Tensor *clamp(Tensor *a, float lo, float hi);
175 Tensor *maximumScalar(Tensor *a, float s);
177 Tensor *minimumScalar(Tensor *a, float s);
178
180 Tensor *matmul(Tensor *a, Tensor *b);
182 Tensor *transpose(Tensor *a);
184 Tensor *permute2(Tensor *a, int a0, int a1);
186 Tensor *permute3(Tensor *a, int a0, int a1, int a2);
188 Tensor *permute4(Tensor *a, int a0, int a1, int a2, int a3);
190 Tensor *permute5(Tensor *a, int a0, int a1, int a2, int a3, int a4);
192 Tensor *permute6(Tensor *a, int a0, int a1, int a2, int a3, int a4, int a5);
194 Tensor *reshape1(Tensor *a, int d0);
196 Tensor *reshape2(Tensor *a, int d0, int d1);
198 Tensor *reshape3(Tensor *a, int d0, int d1, int d2);
200 Tensor *reshape4(Tensor *a, int d0, int d1, int d2, int d3);
202 Tensor *reshape5(Tensor *a, int d0, int d1, int d2, int d3, int d4);
204 Tensor *reshape6(Tensor *a, int d0, int d1, int d2, int d3, int d4, int d5);
206 Tensor *flatten(Tensor *a);
208 Tensor *where(Tensor *cond, Tensor *a, Tensor *b);
210 Tensor *concatN(Tensor *const *ins, int n, int axis);
211
212 // --- neural / speech / terrain ops (eager + traceable) ---
214 Tensor *softmax(Tensor *a, int axis);
216 Tensor *logSoftmax(Tensor *a, int axis);
218 Tensor *layernorm(Tensor *a, float eps);
220 Tensor *layernormWB(Tensor *a, Tensor *scale, Tensor *bias, float eps);
222 Tensor *rmsnorm(Tensor *a, float eps);
224 Tensor *rmsnormW(Tensor *a, Tensor *scale, float eps);
226 Tensor *conv1d(Tensor *x, Tensor *w, int stride, int pad);
228 Tensor *conv1dBias(Tensor *x, Tensor *w, Tensor *bias, int stride, int pad);
230 Tensor *conv2d(Tensor *x, Tensor *w, int stride, int pad);
232 Tensor *conv2dBias(Tensor *x, Tensor *w, Tensor *bias, int stride, int pad);
234 Tensor *maxpool2d(Tensor *x, int ksize, int stride, int pad);
236 Tensor *avgpool2d(Tensor *x, int ksize, int stride, int pad);
238 Tensor *embedding(Tensor *table, Tensor *indices);
240 Tensor *concat2(Tensor *a, Tensor *b, int axis);
242 Tensor *concat3(Tensor *a, Tensor *b, Tensor *c, int axis);
244 Tensor *concat4(Tensor *a, Tensor *b, Tensor *c, Tensor *d, int axis);
246 Tensor *slice(Tensor *a, int axis, int begin, int end);
248 Tensor *sumAxis(Tensor *a, int axis, int keepDims);
250 Tensor *meanAxis(Tensor *a, int axis, int keepDims);
252 Tensor *minAxis(Tensor *a, int axis, int keepDims);
254 Tensor *maxAxis(Tensor *a, int axis, int keepDims);
256 Tensor *argmax(Tensor *a, int axis, int keepDims);
258 Tensor *cast(Tensor *a, const std::string &dtype);
260 Tensor *sdpa(Tensor *q, Tensor *k, Tensor *v, float scale);
262 Tensor *sdpaMasked(Tensor *q, Tensor *k, Tensor *v, Tensor *mask, float scale);
264 Tensor *resize2d(Tensor *a, int outW, int outH, int mode);
265
272 Tensor *quantizeWeight(Tensor *a, const std::string &dtype, int group = 0);
273
275 float reduceSum(Tensor *a);
277 float reduceMean(Tensor *a);
279 float reduceMin(Tensor *a);
281 float reduceMax(Tensor *a);
282
283private:
284 uint32_t seed_ = 1;
285 mutable uint32_t rngState_ = 1;
286 std::vector<Func *> traceStack_;
287
288 float nextUniform() const;
289 float nextGaussian() const;
290 Tensor *filled(const int *dims, int rank, float value);
291};
292
293} // namespace eve::tensor
double value
Duration start
float w
Definition AnimClip.cpp:738
float x
Definition AnimClip.cpp:738
const std::string & s
int mask
building::EdgeCurveGroup group
#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
float f
std::uint32_t seed
Definition PointSet.cpp:807
float begin
float d
glm::vec3 eye
float bias
EVENGINE_API_FOUNDATION public API.
Definition Module.h:46
Trace builder — TF2 tf.function analogue (tf.func in scripts). While active, TF ops record into this ...
Definition Graph.h:137
TF2-like namespace module. Script: tf <- eve.TF(); Default eager; tf.func() traces a graph for compil...
Definition TF.h:20
Tensor * randn4(int d0, int d1, int d2, int d3)
Randn 4.
Definition TF.h:119
Tensor * maxAxis(Tensor *a, int axis, int keepDims)
Max axis.
Tensor * neg(Tensor *a)
Neg.
Tensor * rand3(int d0, int d1, int d2)
Rand 3.
Definition TF.h:106
Tensor * add(Tensor *a, Tensor *b)
Adds add.
Tensor * multiply(Tensor *a, Tensor *b)
Multiply.
Tensor * relu(Tensor *a)
Relu.
Tensor * rand4(int d0, int d1, int d2, int d3)
Rand 4.
Definition TF.h:108
Tensor * tanh(Tensor *a)
Tanh.
Tensor * exp(Tensor *a)
Exp.
Tensor * sigmoid(Tensor *a)
Sigmoid.
Tensor * gelu(Tensor *a)
Gelu.
Tensor * sumAxis(Tensor *a, int axis, int keepDims)
Sum axis.
Tensor * log(Tensor *a)
Log.
Tensor * abs(Tensor *a)
Abs.
Tensor * randn2(int d0, int d1)
Randn 2.
Definition TF.h:115
Tensor * sqrt(Tensor *a)
Sqrt.
Tensor * rand2(int d0, int d1)
Rand 2.
Definition TF.h:104
Tensor * meanAxis(Tensor *a, int axis, int keepDims)
Mean axis.
Tensor * minAxis(Tensor *a, int axis, int keepDims)
Min axis.
Tensor * rand1(int d0)
Rand 1.
Definition TF.h:102
Tensor * sub(Tensor *a, Tensor *b)
Sub.
~TF() override=default
Tf.
Tensor * randn3(int d0, int d1, int d2)
Randn 3.
Definition TF.h:117
Tensor * sin(Tensor *a)
Sin.
Tensor * cos(Tensor *a)
Cos.
Tensor * div(Tensor *a, Tensor *b)
Div.
Tensor * silu(Tensor *a)
Silu.
Tensor * randn1(int d0)
Randn 1.
Definition TF.h:113
float32 / int32 tensor (rank 1–6), row-major. Eager: owns a buffer. Symbolic: node in a Func graph (n...
Definition Tensor.h:48
uint32_t pad[2]