v2.0.0
Loading...
Searching...
No Matches
ml_tensor.h
Go to the documentation of this file.
1//=============================================================================================================
29
30#ifndef ML_TENSOR_H
31#define ML_TENSOR_H
32
33//=============================================================================================================
34// INCLUDES
35//=============================================================================================================
36
37#include "ml_global.h"
38
39//=============================================================================================================
40// EIGEN INCLUDES
41//=============================================================================================================
42
43#include <Eigen/Core>
44
45//=============================================================================================================
46// STL INCLUDES
47//=============================================================================================================
48
49#include <cassert>
50#include <cstdint>
51#include <memory>
52#include <vector>
53
54//=============================================================================================================
55// DEFINE NAMESPACE MLLIB
56//=============================================================================================================
57
58namespace MLLIB
59{
60
61//=============================================================================================================
76{
77public:
78 // --- type aliases used in the public API --------------------------------
79 using RowMajorMatrixXf = Eigen::Matrix<float, Eigen::Dynamic, Eigen::Dynamic, Eigen::RowMajor>;
80 using RowMajorMatrixMap = Eigen::Map<RowMajorMatrixXf>;
81 using ConstRowMajorMatrixMap = Eigen::Map<const RowMajorMatrixXf>;
82
83 //=========================================================================================================
87 MlTensor();
88
89 //=========================================================================================================
97 MlTensor(std::vector<float>&& data, std::vector<int64_t> shape);
98
99 //=========================================================================================================
106 MlTensor(const float* data, std::vector<int64_t> shape);
107
108 //=========================================================================================================
116 explicit MlTensor(const Eigen::MatrixXf& mat);
117
118 //=========================================================================================================
126 explicit MlTensor(const Eigen::MatrixXd& mat);
127
128 //=========================================================================================================
138 static MlTensor view(float* data, std::vector<int64_t> shape);
139
140 //=========================================================================================================
149 static MlTensor fromBuffer(const float* data, int rows, int cols);
150
151 // --- shape access -------------------------------------------------------
152
153 //=========================================================================================================
157 int ndim() const;
158
159 //=========================================================================================================
163 int64_t size() const;
164
165 //=========================================================================================================
169 const std::vector<int64_t>& shape() const;
170
171 //=========================================================================================================
176 int64_t shape(int dim) const;
177
178 //=========================================================================================================
183 int rows() const;
184
185 //=========================================================================================================
190 int cols() const;
191
192 // --- raw data access (zero-copy) ----------------------------------------
193
194 //=========================================================================================================
198 float* data();
199
200 //=========================================================================================================
204 const float* data() const;
205
206 // --- Eigen Map accessors (zero-copy, 2-D) -------------------------------
207
208 //=========================================================================================================
216
217 //=========================================================================================================
225
226 // --- Eigen copy helpers (produce column-major copies) -------------------
227
228 //=========================================================================================================
232 Eigen::MatrixXf toMatrixXf() const;
233
234 //=========================================================================================================
238 Eigen::MatrixXd toMatrixXd() const;
239
240 // --- reshape / query ----------------------------------------------------
241
242 //=========================================================================================================
252 MlTensor reshape(std::vector<int64_t> newShape) const;
253
254 //=========================================================================================================
258 bool isView() const;
259
260 //=========================================================================================================
264 bool empty() const;
265
266private:
267 static int64_t computeSize(const std::vector<int64_t>& shape);
268
269 std::shared_ptr<std::vector<float>> m_storage;
270 float* m_data = nullptr;
271 std::vector<int64_t> m_shape;
272 int64_t m_size = 0;
273};
274
275} // namespace MLLIB
276
277#endif // ML_TENSOR_H
Export/import macros, build-stamp accessors and namespace anchor for the MLLIB machine-learning libra...
#define MLSHARED_EXPORT
Definition ml_global.h:48
Tensors, model abstraction, ONNX Runtime inference and Python training drivers used across mne-cpp.
int cols() const
Eigen::MatrixXf toMatrixXf() const
Eigen::Map< RowMajorMatrixXf > RowMajorMatrixMap
Definition ml_tensor.h:80
Eigen::Matrix< float, Eigen::Dynamic, Eigen::Dynamic, Eigen::RowMajor > RowMajorMatrixXf
Definition ml_tensor.h:79
static MlTensor fromBuffer(const float *data, int rows, int cols)
Eigen::MatrixXd toMatrixXd() const
static MlTensor view(float *data, std::vector< int64_t > shape)
MlTensor(const Eigen::MatrixXd &mat)
int ndim() const
RowMajorMatrixMap matrix()
int64_t size() const
MlTensor(const Eigen::MatrixXf &mat)
int rows() const
MlTensor reshape(std::vector< int64_t > newShape) const
bool empty() const
const std::vector< int64_t > & shape() const
bool isView() const
Eigen::Map< const RowMajorMatrixXf > ConstRowMajorMatrixMap
Definition ml_tensor.h:81