载入中...
搜索中...
未找到
Tensor.h
浏览该文件的文档.
1
2#include "common/Export.h"
3#ifndef EVE_TENSOR_TENSOR_H
4#define EVE_TENSOR_TENSOR_H
5
6#include <cstdint>
7#include <string>
8#include <vector>
9
10namespace eve::tensor {
11
12class Graph;
13class Func;
14class TF;
15
24enum class DType : uint8_t {
25 Float32 = 0,
26 Int32 = 1,
27 Fp16 = 2, // IEEE half, 2 bytes/elem, weight-only
28 Fp8E4M3 = 3, // 1 byte/elem, weight-only
29 Fp4E2M1 = 4, // 4-bit e2m1, two per byte, weight-only
30 Int8 = 5, // 1 byte/elem + per-group scale, weight-only
31 Int4 = 6, // 4-bit + per-group scale, weight-only
32};
33
34namespace q {
36bool isQuantDType(DType dt);
37}
38
40const char *dtypeName(DType dtype);
42bool parseDType(const std::string &name, DType &out);
43
49public:
50 static constexpr int kMaxRank = 6;
51
53 Tensor() = default;
55 explicit Tensor(const int *dims, int rank);
57 explicit Tensor(DType dtype, const int *dims, int rank);
59 Tensor(int d0);
61 Tensor(int d0, int d1);
63 Tensor(int d0, int d1, int d2);
65 Tensor(int d0, int d1, int d2, int d3);
67 Tensor(int d0, int d1, int d2, int d3, int d4);
69 Tensor(int d0, int d1, int d2, int d3, int d4, int d5);
70
72 static Tensor *makeSymbolic(Graph *graph, int nodeId, const int *dims, int rank);
73
75 bool isSymbolic() const { return kind_ == Kind::Symbolic; }
77 bool isEager() const { return kind_ == Kind::Eager; }
79 Graph *graph() const { return graph_; }
81 int nodeId() const { return nodeId_; }
82
84 int getRank() const { return rank_; }
86 int getSize() const { return size_; }
88 int getDim(int axis) const;
90 int getDim0() const { return dims_[0]; }
92 int getDim1() const { return dims_[1]; }
94 int getDim2() const { return dims_[2]; }
96 int getDim3() const { return dims_[3]; }
98 int getDim4() const { return dims_[4]; }
100 int getDim5() const { return dims_[5]; }
102 std::string getDevice() const { return device_; }
104 std::string getDtype() const { return dtypeName(dtype_); }
106 DType dtype() const { return dtype_; }
108 void setDtype(DType dtype) { dtype_ = dtype; }
109
112 bool isQuantized() const { return q::isQuantDType(dtype_); }
113
116 std::vector<float> dequantized() const;
117
120 const std::vector<float> &qScales() const { return qScales_; }
122 int qGroup() const { return qGroup_; }
124 const std::vector<uint8_t> &qBytes() const { return bytes_; }
125
127 float get(int flatIndex) const;
129 void set(int flatIndex, float value);
131 float get1(int i0) const;
133 void set1(int i0, float value);
135 float get2(int i0, int i1) const;
137 void set2(int i0, int i1, float value);
139 float get3(int i0, int i1, int i2) const;
141 void set3(int i0, int i1, int i2, float value);
143 float get4(int i0, int i1, int i2, int i3) const;
145 void set4(int i0, int i1, int i2, int i3, float value);
147 float get5(int i0, int i1, int i2, int i3, int i4) const;
149 void set5(int i0, int i1, int i2, int i3, int i4, float value);
151 float get6(int i0, int i1, int i2, int i3, int i4, int i5) const;
153 void set6(int i0, int i1, int i2, int i3, int i4, int i5, float value);
154
156 void fill(float value);
158 void copyFrom(const Tensor *other);
160 Tensor *clone() const;
161
163 Tensor *add(const Tensor *other) const;
165 Tensor *sub(const Tensor *other) const;
167 Tensor *multiply(const Tensor *other) const;
169 Tensor *div(const Tensor *other) const;
171 Tensor *addScalar(float s) const;
173 Tensor *subScalar(float s) const;
175 Tensor *mulScalar(float s) const;
177 Tensor *divScalar(float s) const;
179 Tensor *neg() const;
181 Tensor *abs() const;
183 Tensor *sqrt() const;
185 Tensor *exp() const;
187 Tensor *log() const;
189 Tensor *sin() const;
191 Tensor *cos() const;
193 Tensor *tanh() const;
195 Tensor *relu() const;
197 Tensor *sigmoid() const;
199 Tensor *gelu() const;
201 Tensor *silu() const;
203 Tensor *powScalar(float exp) const;
205 Tensor *clamp(float lo, float hi) const;
207 Tensor *maximumScalar(float s) const;
209 Tensor *minimumScalar(float s) const;
210
212 void addInPlace(const Tensor *other);
214 void multiplyInPlace(const Tensor *other);
216 void addScalarInPlace(float s);
218 void mulScalarInPlace(float s);
220 void reluInPlace();
221
223 float reduceSum() const;
225 float reduceMean() const;
227 float reduceMin() const;
229 float reduceMax() const;
231 float dot(const Tensor *other) const;
232
234 Tensor *matmul(const Tensor *other) const;
236 Tensor *transpose() const;
238 Tensor *permute(const int *order, int rank) const;
240 Tensor *reshape1(int d0) const;
242 Tensor *reshape2(int d0, int d1) const;
244 Tensor *reshape3(int d0, int d1, int d2) const;
246 Tensor *reshape4(int d0, int d1, int d2, int d3) const;
248 Tensor *reshape5(int d0, int d1, int d2, int d3, int d4) const;
250 Tensor *reshape6(int d0, int d1, int d2, int d3, int d4, int d5) const;
252 Tensor *flatten() const;
253
255 float *data();
257 const float *data() const;
258
260 void ensureEager(const char *op) const;
261
263 static int product(const int *dims, int rank);
264
265private:
266 friend class TF;
267 friend class Func;
268 friend class CompiledFunction;
269 friend class Graph;
270
271 enum class Kind { Eager, Symbolic };
272
273 void initDims(DType dtype, const int *dims, int rank);
274 void checkSameShape(const Tensor *other, const char *op) const;
275 int offset2(int i0, int i1) const;
276 int offset3(int i0, int i1, int i2) const;
277 int offset4(int i0, int i1, int i2, int i3) const;
278 int offset5(int i0, int i1, int i2, int i3, int i4) const;
279 int offset6(int i0, int i1, int i2, int i3, int i4, int i5) const;
280
281 Kind kind_ = Kind::Eager;
282 Graph *graph_ = nullptr;
283 int nodeId_ = -1;
284 int rank_ = 0;
285 int dims_[kMaxRank] = {0, 0, 0, 0, 0, 0};
286 int size_ = 0;
287 DType dtype_ = DType::Float32;
288 std::vector<float> data_;
289 std::vector<uint8_t> bytes_; // packed payload for quantized dtypes
290 std::vector<float> qScales_; // per-group scales (int8/int4)
291 int qGroup_ = 0; // elements per scale group
292 std::string device_ = "cpu";
293};
294
295} // namespace eve::tensor
296
297#endif // EVE_TENSOR_TENSOR_H
double value
const std::string & s
std::string nodeId
#define EVENGINE_API_DOMAINS
Definition Export.h:110
uint32_t i1
Definition Grass.cpp:61
uint32_t i2
Definition Grass.cpp:61
uint32_t i0
Definition Grass.cpp:61
std::array< double, 10 > q
std::string name
std::vector< std::int32_t > order
std::map< std::string, std::vector< std::string > > graph
Definition Package.cpp:59
Optimized / scheduled graph ready to run with feeds.
Definition Graph.h:242
Trace builder — TF2 tf.function analogue (tf.func in scripts). While active, TF ops record into this ...
Definition Graph.h:137
EVENGINE_API_DOMAINS public API.
Definition Graph.h:111
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
Tensor * tanh() const
Tanh.
const std::vector< float > & qScales() const
Q scales.
Definition Tensor.h:120
std::string getDtype() const
Returns the dtype.
Definition Tensor.h:104
Tensor * sin() const
Sin.
int getDim4() const
Returns the dim 4.
Definition Tensor.h:98
int qGroup() const
Q group.
Definition Tensor.h:122
bool isQuantized() const
True when quantized.
Definition Tensor.h:112
Tensor * log() const
Log.
int getDim2() const
Returns the dim 2.
Definition Tensor.h:94
Tensor * sigmoid() const
Sigmoid.
int getDim3() const
Returns the dim 3.
Definition Tensor.h:96
Tensor * abs() const
Abs.
const std::vector< uint8_t > & qBytes() const
Q bytes.
Definition Tensor.h:124
std::string getDevice() const
Returns the device.
Definition Tensor.h:102
int getSize() const
Byte length of the owned buffer.
Definition Tensor.h:86
Tensor * neg() const
Neg.
Tensor * relu() const
Relu.
Tensor()=default
Tensor.
int nodeId() const
Node id.
Definition Tensor.h:81
Tensor * exp() const
Exp.
Tensor * gelu() const
Gelu.
bool isSymbolic() const
True when symbolic.
Definition Tensor.h:75
int getDim1() const
Returns the dim 1.
Definition Tensor.h:92
int getDim0() const
Returns the dim 0.
Definition Tensor.h:90
Graph * graph() const
Graph.
Definition Tensor.h:79
int getRank() const
Returns the rank.
Definition Tensor.h:84
Tensor * sqrt() const
Sqrt.
Tensor * silu() const
Silu.
bool isEager() const
True when eager.
Definition Tensor.h:77
DType dtype() const
Dtype.
Definition Tensor.h:106
void setDtype(DType dtype)
Sets the dtype.
Definition Tensor.h:108
int getDim5() const
Returns the dim 5.
Definition Tensor.h:100
Tensor * cos() const
Cos.
bool isQuantDType(DType dt)
True when quant d type.
Definition Quant.h:17
bool parseDType(const std::string &name, DType &out)
Parse d type.
Definition Tensor.cpp:39
DType
Tensor element types.
Definition Tensor.h:24
const char * dtypeName(DType dtype)
Dtype name.
Definition Tensor.cpp:26