Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorIO.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2014 Benoit Steiner <benoit.steiner.goog@gmail.com>
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9// SPDX-License-Identifier: MPL-2.0
10
11#ifndef EIGEN_TENSOR_TENSOR_IO_H
12#define EIGEN_TENSOR_TENSOR_IO_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19struct TensorIOFormat;
20
21namespace internal {
22template <typename Tensor, std::size_t rank, typename Format, typename EnableIf = void>
23struct TensorPrinter;
24}
25
26template <typename Derived_>
27struct TensorIOFormatBase {
28 using Derived = Derived_;
29 TensorIOFormatBase(const std::vector<std::string>& separator, const std::vector<std::string>& prefix,
30 const std::vector<std::string>& suffix, int precision = StreamPrecision, int flags = 0,
31 const std::string& tenPrefix = "", const std::string& tenSuffix = "", const char fill = ' ')
32 : tenPrefix(tenPrefix),
33 tenSuffix(tenSuffix),
34 prefix(prefix),
35 suffix(suffix),
36 separator(separator),
37 fill(fill),
38 precision(precision),
39 flags(flags) {
40 init_spacer();
41 }
42
43 void init_spacer() {
44 if (flags & DontAlignCols) return;
45 spacer.resize(prefix.size());
46 spacer[0] = "";
47 int i = int(tenPrefix.length()) - 1;
48 while (i >= 0 && tenPrefix[i] != '\n') {
49 spacer[0] += ' ';
50 i--;
51 }
52
53 for (std::size_t k = 1; k < prefix.size(); k++) {
54 int j = int(prefix[k].length()) - 1;
55 while (j >= 0 && prefix[k][j] != '\n') {
56 spacer[k] += ' ';
57 j--;
58 }
59 }
60 }
61
62 std::string tenPrefix;
63 std::string tenSuffix;
64 std::vector<std::string> prefix;
65 std::vector<std::string> suffix;
66 std::vector<std::string> separator;
67 char fill;
68 int precision;
69 int flags;
70 std::vector<std::string> spacer{};
71};
72
73struct TensorIOFormatNumpy : public TensorIOFormatBase<TensorIOFormatNumpy> {
74 using Base = TensorIOFormatBase<TensorIOFormatNumpy>;
75 TensorIOFormatNumpy()
76 : Base(/*separator=*/{" ", "\n"}, /*prefix=*/{"", "["}, /*suffix=*/{"", "]"}, /*precision=*/StreamPrecision,
77 /*flags=*/0, /*tenPrefix=*/"[", /*tenSuffix=*/"]") {}
78};
79
80struct TensorIOFormatNative : public TensorIOFormatBase<TensorIOFormatNative> {
81 using Base = TensorIOFormatBase<TensorIOFormatNative>;
82 TensorIOFormatNative()
83 : Base(/*separator=*/{", ", ",\n", "\n"}, /*prefix=*/{"", "{"}, /*suffix=*/{"", "}"},
84 /*precision=*/StreamPrecision, /*flags=*/0, /*tenPrefix=*/"{", /*tenSuffix=*/"}") {}
85};
86
87struct TensorIOFormatPlain : public TensorIOFormatBase<TensorIOFormatPlain> {
88 using Base = TensorIOFormatBase<TensorIOFormatPlain>;
89 TensorIOFormatPlain()
90 : Base(/*separator=*/{" ", "\n", "\n", ""}, /*prefix=*/{""}, /*suffix=*/{""}, /*precision=*/StreamPrecision,
91 /*flags=*/0, /*tenPrefix=*/"", /*tenSuffix=*/"") {}
92};
93
94struct TensorIOFormatLegacy : public TensorIOFormatBase<TensorIOFormatLegacy> {
95 using Base = TensorIOFormatBase<TensorIOFormatLegacy>;
96 TensorIOFormatLegacy()
97 : Base(/*separator=*/{", ", "\n"}, /*prefix=*/{"", "["}, /*suffix=*/{"", "]"}, /*precision=*/StreamPrecision,
98 /*flags=*/0, /*tenPrefix=*/"", /*tenSuffix=*/"") {}
99};
100
101struct TensorIOFormat : public TensorIOFormatBase<TensorIOFormat> {
102 using Base = TensorIOFormatBase<TensorIOFormat>;
103 TensorIOFormat(const std::vector<std::string>& separator, const std::vector<std::string>& prefix,
104 const std::vector<std::string>& suffix, int precision = StreamPrecision, int flags = 0,
105 const std::string& tenPrefix = "", const std::string& tenSuffix = "", const char fill = ' ')
106 : Base(separator, prefix, suffix, precision, flags, tenPrefix, tenSuffix, fill) {}
107
108 static inline const TensorIOFormatNumpy Numpy() { return TensorIOFormatNumpy{}; }
109
110 static inline const TensorIOFormatPlain Plain() { return TensorIOFormatPlain{}; }
111
112 static inline const TensorIOFormatNative Native() { return TensorIOFormatNative{}; }
113
114 static inline const TensorIOFormatLegacy Legacy() { return TensorIOFormatLegacy{}; }
115};
116
117template <typename T, int Layout, int rank, typename Format>
118class TensorWithFormat;
119// specialize for Layout=ColMajor, Layout=RowMajor and rank=0.
120template <typename T, int rank, typename Format>
121class TensorWithFormat<T, RowMajor, rank, Format> {
122 public:
123 TensorWithFormat(const T& tensor, const Format& format) : t_tensor(tensor), t_format(format) {}
124
125 friend std::ostream& operator<<(std::ostream& os, const TensorWithFormat<T, RowMajor, rank, Format>& wf) {
126 // Evaluate the expression if needed
127 typedef TensorEvaluator<const TensorForcedEvalOp<const T>, DefaultDevice> Evaluator;
128 TensorForcedEvalOp<const T> eval = wf.t_tensor.eval();
129 Evaluator tensor(eval, DefaultDevice());
130 tensor.evalSubExprsIfNeeded(nullptr);
131 internal::TensorPrinter<Evaluator, rank, Format>::run(os, tensor, wf.t_format);
132 // Cleanup.
133 tensor.cleanup();
134 return os;
135 }
136
137 protected:
138 T t_tensor;
139 Format t_format;
140};
141
142template <typename T, int rank, typename Format>
143class TensorWithFormat<T, ColMajor, rank, Format> {
144 public:
145 TensorWithFormat(const T& tensor, const Format& format) : t_tensor(tensor), t_format(format) {}
146
147 friend std::ostream& operator<<(std::ostream& os, const TensorWithFormat<T, ColMajor, rank, Format>& wf) {
148 // Switch to RowMajor storage and print afterwards
149 typedef typename T::Index IndexType;
150 std::array<IndexType, rank> shuffle;
151 std::array<IndexType, rank> id;
152 std::iota(id.begin(), id.end(), IndexType(0));
153 std::copy(id.begin(), id.end(), shuffle.rbegin());
154 auto tensor_row_major = wf.t_tensor.swap_layout().shuffle(shuffle);
155
156 // Evaluate the expression if needed
157 typedef TensorEvaluator<const TensorForcedEvalOp<const decltype(tensor_row_major)>, DefaultDevice> Evaluator;
158 TensorForcedEvalOp<const decltype(tensor_row_major)> eval = tensor_row_major.eval();
159 Evaluator tensor(eval, DefaultDevice());
160 tensor.evalSubExprsIfNeeded(nullptr);
161 internal::TensorPrinter<Evaluator, rank, Format>::run(os, tensor, wf.t_format);
162 // Cleanup.
163 tensor.cleanup();
164 return os;
165 }
166
167 protected:
168 T t_tensor;
169 Format t_format;
170};
171
172template <typename T, typename Format>
173class TensorWithFormat<T, ColMajor, 0, Format> {
174 public:
175 TensorWithFormat(const T& tensor, const Format& format) : t_tensor(tensor), t_format(format) {}
176
177 friend std::ostream& operator<<(std::ostream& os, const TensorWithFormat<T, ColMajor, 0, Format>& wf) {
178 // Evaluate the expression if needed
179 typedef TensorEvaluator<const TensorForcedEvalOp<const T>, DefaultDevice> Evaluator;
180 TensorForcedEvalOp<const T> eval = wf.t_tensor.eval();
181 Evaluator tensor(eval, DefaultDevice());
182 tensor.evalSubExprsIfNeeded(nullptr);
183 internal::TensorPrinter<Evaluator, 0, Format>::run(os, tensor, wf.t_format);
184 // Cleanup.
185 tensor.cleanup();
186 return os;
187 }
188
189 protected:
190 T t_tensor;
191 Format t_format;
192};
193
194namespace internal {
195
196// Default scalar printer.
197template <typename Scalar, typename Format, typename EnableIf = void>
198struct ScalarPrinter {
199 static void run(std::ostream& stream, const Scalar& scalar, const Format&) { stream << scalar; }
200};
201
202template <typename Scalar>
203struct ScalarPrinter<Scalar, TensorIOFormatNumpy, std::enable_if_t<NumTraits<Scalar>::IsComplex>> {
204 static void run(std::ostream& stream, const Scalar& scalar, const TensorIOFormatNumpy&) {
205 stream << numext::real(scalar) << "+" << numext::imag(scalar) << "j";
206 }
207};
208
209template <typename Scalar>
210struct ScalarPrinter<Scalar, TensorIOFormatNative, std::enable_if_t<NumTraits<Scalar>::IsComplex>> {
211 static void run(std::ostream& stream, const Scalar& scalar, const TensorIOFormatNative&) {
212 stream << "{" << numext::real(scalar) << ", " << numext::imag(scalar) << "}";
213 }
214};
215
216template <typename Tensor, std::size_t rank, typename Format, typename EnableIf>
217struct TensorPrinter {
218 using Scalar = std::remove_const_t<typename Tensor::Scalar>;
219
220 static void run(std::ostream& s, const Tensor& tensor, const Format& fmt) {
221 typedef typename Tensor::Index IndexType;
222
223 eigen_assert(Tensor::Layout == RowMajor);
224 typedef std::conditional_t<std::is_same<Scalar, char>::value || std::is_same<Scalar, unsigned char>::value ||
225 std::is_same<Scalar, numext::int8_t>::value ||
226 std::is_same<Scalar, numext::uint8_t>::value,
227 int,
228 std::conditional_t<std::is_same<Scalar, std::complex<char>>::value ||
229 std::is_same<Scalar, std::complex<unsigned char>>::value ||
230 std::is_same<Scalar, std::complex<numext::int8_t>>::value ||
231 std::is_same<Scalar, std::complex<numext::uint8_t>>::value,
232 std::complex<int>, const Scalar&>>
233 PrintType;
234
235 const IndexType total_size = array_prod(tensor.dimensions());
236
237 std::streamsize explicit_precision;
238 if (fmt.precision == StreamPrecision) {
239 explicit_precision = 0;
240 } else if (fmt.precision == FullPrecision) {
241 EIGEN_IF_CONSTEXPR (NumTraits<Scalar>::IsInteger) {
242 explicit_precision = 0;
243 } else {
244 explicit_precision = significant_decimals_impl<Scalar>::run();
245 }
246 } else {
247 explicit_precision = fmt.precision;
248 }
249
250 std::streamsize old_precision = 0;
251 if (explicit_precision) old_precision = s.precision(explicit_precision);
252
253 IndexType width = 0;
254 bool align_cols = !(fmt.flags & DontAlignCols);
255 if (align_cols) {
256 // compute the largest width
257 for (IndexType i = 0; i < total_size; i++) {
258 std::stringstream sstr;
259 sstr.copyfmt(s);
260 ScalarPrinter<Scalar, Format>::run(sstr, static_cast<PrintType>(tensor.data()[i]), fmt);
261 width = std::max<IndexType>(width, IndexType(sstr.str().length()));
262 }
263 }
264 s << fmt.tenPrefix;
265 for (IndexType i = 0; i < total_size; i++) {
266 std::array<bool, rank> is_at_end{};
267 std::array<bool, rank> is_at_begin{};
268
269 // is the ith element the end of a coeff (always true), of a row, of a matrix, ...?
270 for (std::size_t k = 0; k < rank; k++) {
271 if ((i + 1) % (std::accumulate(tensor.dimensions().rbegin(), tensor.dimensions().rbegin() + k, 1,
272 std::multiplies<IndexType>())) ==
273 0) {
274 is_at_end[k] = true;
275 }
276 }
277
278 // is the ith element the begin of a coeff (always true), of a row, of a matrix, ...?
279 for (std::size_t k = 0; k < rank; k++) {
280 if (i % (std::accumulate(tensor.dimensions().rbegin(), tensor.dimensions().rbegin() + k, 1,
281 std::multiplies<IndexType>())) ==
282 0) {
283 is_at_begin[k] = true;
284 }
285 }
286
287 // do we have a line break?
288 bool is_at_begin_after_newline = false;
289 for (std::size_t k = 0; k < rank; k++) {
290 if (is_at_begin[k]) {
291 std::size_t separator_index = (k < fmt.separator.size()) ? k : fmt.separator.size() - 1;
292 if (fmt.separator[separator_index].find('\n') != std::string::npos) {
293 is_at_begin_after_newline = true;
294 }
295 }
296 }
297
298 bool is_at_end_before_newline = false;
299 for (std::size_t k = 0; k < rank; k++) {
300 if (is_at_end[k]) {
301 std::size_t separator_index = (k < fmt.separator.size()) ? k : fmt.separator.size() - 1;
302 if (fmt.separator[separator_index].find('\n') != std::string::npos) {
303 is_at_end_before_newline = true;
304 }
305 }
306 }
307
308 std::stringstream suffix, prefix, separator;
309 for (std::size_t k = 0; k < rank; k++) {
310 std::size_t suffix_index = (k < fmt.suffix.size()) ? k : fmt.suffix.size() - 1;
311 if (is_at_end[k]) {
312 suffix << fmt.suffix[suffix_index];
313 }
314 }
315 for (std::size_t k = 0; k < rank; k++) {
316 std::size_t separator_index = (k < fmt.separator.size()) ? k : fmt.separator.size() - 1;
317 if (is_at_end[k] &&
318 (!is_at_end_before_newline || fmt.separator[separator_index].find('\n') != std::string::npos)) {
319 separator << fmt.separator[separator_index];
320 }
321 }
322 for (std::size_t k = 0; k < rank; k++) {
323 std::size_t spacer_index = (k < fmt.spacer.size()) ? k : fmt.spacer.size() - 1;
324 if (i != 0 && is_at_begin_after_newline && (!is_at_begin[k] || k == 0)) {
325 prefix << fmt.spacer[spacer_index];
326 }
327 }
328 for (int k = rank - 1; k >= 0; k--) {
329 std::size_t prefix_index = (static_cast<std::size_t>(k) < fmt.prefix.size()) ? k : fmt.prefix.size() - 1;
330 if (is_at_begin[k]) {
331 prefix << fmt.prefix[prefix_index];
332 }
333 }
334
335 s << prefix.str();
336 // So we don't mess around with formatting, output scalar to a string stream, and adjust the width/fill manually.
337 std::stringstream sstr;
338 sstr.copyfmt(s);
339 ScalarPrinter<Scalar, Format>::run(sstr, static_cast<PrintType>(tensor.data()[i]), fmt);
340 std::string scalar_str = sstr.str();
341 IndexType scalar_width = scalar_str.length();
342 if (width && scalar_width < width) {
343 std::string filler;
344 for (IndexType j = scalar_width; j < width; ++j) {
345 filler.push_back(fmt.fill);
346 }
347 s << filler;
348 }
349 s << scalar_str;
350 s << suffix.str();
351 if (i < total_size - 1) {
352 s << separator.str();
353 }
354 }
355 s << fmt.tenSuffix;
356 if (explicit_precision) s.precision(old_precision);
357 }
358};
359
360template <typename Tensor, std::size_t rank>
361struct TensorPrinter<Tensor, rank, TensorIOFormatLegacy, std::enable_if_t<rank != 0>> {
362 using Format = TensorIOFormatLegacy;
363 using Scalar = std::remove_const_t<typename Tensor::Scalar>;
364
365 static void run(std::ostream& s, const Tensor& tensor, const Format&) {
366 typedef typename Tensor::Index IndexType;
367 // backwards compatibility case: print tensor after reshaping to matrix of size dim(0) x
368 // (dim(1)*dim(2)*...*dim(rank-1)).
369 const IndexType total_size = internal::array_prod(tensor.dimensions());
370 if (total_size > 0) {
371 const IndexType first_dim = Eigen::internal::array_get<0>(tensor.dimensions());
372 Map<const Array<Scalar, Dynamic, Dynamic, Tensor::Layout>> matrix(tensor.data(), first_dim,
373 total_size / first_dim);
374 s << matrix;
375 return;
376 }
377 }
378};
379
380template <typename Tensor, typename Format>
381struct TensorPrinter<Tensor, 0, Format> {
382 static void run(std::ostream& s, const Tensor& tensor, const Format& fmt) {
383 using Scalar = std::remove_const_t<typename Tensor::Scalar>;
384
385 std::streamsize explicit_precision;
386 if (fmt.precision == StreamPrecision) {
387 explicit_precision = 0;
388 } else if (fmt.precision == FullPrecision) {
389 EIGEN_IF_CONSTEXPR (NumTraits<Scalar>::IsInteger) {
390 explicit_precision = 0;
391 } else {
392 explicit_precision = significant_decimals_impl<Scalar>::run();
393 }
394 } else {
395 explicit_precision = fmt.precision;
396 }
397
398 std::streamsize old_precision = 0;
399 if (explicit_precision) old_precision = s.precision(explicit_precision);
400 s << fmt.tenPrefix;
401 ScalarPrinter<Scalar, Format>::run(s, tensor.coeff(0), fmt);
402 s << fmt.tenSuffix;
403 if (explicit_precision) s.precision(old_precision);
404 }
405};
406
407} // end namespace internal
408template <typename T>
409std::ostream& operator<<(std::ostream& s, const TensorBase<T, ReadOnlyAccessors>& t) {
410 s << t.format(TensorIOFormat::Plain());
411 return s;
412}
413} // end namespace Eigen
414
415#endif // EIGEN_TENSOR_TENSOR_IO_H
The tensor base class.
Definition TensorForwardDeclarations.h:69
Namespace containing all symbols from the Eigen library.