11#ifndef EIGEN_TENSOR_TENSOR_IO_H
12#define EIGEN_TENSOR_TENSOR_IO_H
15#include "./InternalHeaderCheck.h"
22template <
typename Tensor, std::
size_t rank,
typename Format,
typename EnableIf =
void>
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),
44 if (flags & DontAlignCols)
return;
45 spacer.resize(prefix.size());
47 int i = int(tenPrefix.length()) - 1;
48 while (i >= 0 && tenPrefix[i] !=
'\n') {
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') {
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;
70 std::vector<std::string> spacer{};
73struct TensorIOFormatNumpy :
public TensorIOFormatBase<TensorIOFormatNumpy> {
74 using Base = TensorIOFormatBase<TensorIOFormatNumpy>;
76 : Base({
" ",
"\n"}, {
"",
"["}, {
"",
"]"}, StreamPrecision,
80struct TensorIOFormatNative :
public TensorIOFormatBase<TensorIOFormatNative> {
81 using Base = TensorIOFormatBase<TensorIOFormatNative>;
82 TensorIOFormatNative()
83 : Base({
", ",
",\n",
"\n"}, {
"",
"{"}, {
"",
"}"},
84 StreamPrecision, 0,
"{",
"}") {}
87struct TensorIOFormatPlain :
public TensorIOFormatBase<TensorIOFormatPlain> {
88 using Base = TensorIOFormatBase<TensorIOFormatPlain>;
90 : Base({
" ",
"\n",
"\n",
""}, {
""}, {
""}, StreamPrecision,
94struct TensorIOFormatLegacy :
public TensorIOFormatBase<TensorIOFormatLegacy> {
95 using Base = TensorIOFormatBase<TensorIOFormatLegacy>;
96 TensorIOFormatLegacy()
97 : Base({
", ",
"\n"}, {
"",
"["}, {
"",
"]"}, StreamPrecision,
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) {}
108 static inline const TensorIOFormatNumpy Numpy() {
return TensorIOFormatNumpy{}; }
110 static inline const TensorIOFormatPlain Plain() {
return TensorIOFormatPlain{}; }
112 static inline const TensorIOFormatNative Native() {
return TensorIOFormatNative{}; }
114 static inline const TensorIOFormatLegacy Legacy() {
return TensorIOFormatLegacy{}; }
117template <
typename T,
int Layout,
int rank,
typename Format>
118class TensorWithFormat;
120template <
typename T,
int rank,
typename Format>
121class TensorWithFormat<T,
RowMajor, rank, Format> {
123 TensorWithFormat(
const T& tensor,
const Format& format) : t_tensor(tensor), t_format(format) {}
125 friend std::ostream& operator<<(std::ostream& os,
const TensorWithFormat<T, RowMajor, rank, Format>& wf) {
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);
142template <
typename T,
int rank,
typename Format>
143class TensorWithFormat<T,
ColMajor, rank, Format> {
145 TensorWithFormat(
const T& tensor,
const Format& format) : t_tensor(tensor), t_format(format) {}
147 friend std::ostream& operator<<(std::ostream& os,
const TensorWithFormat<T, ColMajor, rank, Format>& wf) {
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);
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);
172template <
typename T,
typename Format>
173class TensorWithFormat<T,
ColMajor, 0, Format> {
175 TensorWithFormat(
const T& tensor,
const Format& format) : t_tensor(tensor), t_format(format) {}
177 friend std::ostream& operator<<(std::ostream& os,
const TensorWithFormat<T, ColMajor, 0, Format>& wf) {
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);
197template <
typename Scalar,
typename Format,
typename EnableIf =
void>
198struct ScalarPrinter {
199 static void run(std::ostream& stream,
const Scalar& scalar,
const Format&) { stream << scalar; }
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";
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) <<
"}";
216template <
typename Tensor, std::
size_t rank,
typename Format,
typename EnableIf>
217struct TensorPrinter {
218 using Scalar = std::remove_const_t<typename Tensor::Scalar>;
220 static void run(std::ostream& s,
const Tensor& tensor,
const Format& fmt) {
221 typedef typename Tensor::Index IndexType;
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,
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&>>
235 const IndexType total_size = array_prod(tensor.dimensions());
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;
244 explicit_precision = significant_decimals_impl<Scalar>::run();
247 explicit_precision = fmt.precision;
250 std::streamsize old_precision = 0;
251 if (explicit_precision) old_precision = s.precision(explicit_precision);
254 bool align_cols = !(fmt.flags & DontAlignCols);
257 for (IndexType i = 0; i < total_size; i++) {
258 std::stringstream sstr;
260 ScalarPrinter<Scalar, Format>::run(sstr,
static_cast<PrintType
>(tensor.data()[i]), fmt);
261 width = std::max<IndexType>(width, IndexType(sstr.str().length()));
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{};
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>())) ==
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>())) ==
283 is_at_begin[k] =
true;
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;
298 bool is_at_end_before_newline =
false;
299 for (std::size_t k = 0; k < rank; 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;
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;
312 suffix << fmt.suffix[suffix_index];
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;
318 (!is_at_end_before_newline || fmt.separator[separator_index].find(
'\n') != std::string::npos)) {
319 separator << fmt.separator[separator_index];
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];
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];
337 std::stringstream sstr;
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) {
344 for (IndexType j = scalar_width; j < width; ++j) {
345 filler.push_back(fmt.fill);
351 if (i < total_size - 1) {
352 s << separator.str();
356 if (explicit_precision) s.precision(old_precision);
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>;
365 static void run(std::ostream& s,
const Tensor& tensor,
const Format&) {
366 typedef typename Tensor::Index IndexType;
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);
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>;
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;
392 explicit_precision = significant_decimals_impl<Scalar>::run();
395 explicit_precision = fmt.precision;
398 std::streamsize old_precision = 0;
399 if (explicit_precision) old_precision = s.precision(explicit_precision);
401 ScalarPrinter<Scalar, Format>::run(s, tensor.coeff(0), fmt);
403 if (explicit_precision) s.precision(old_precision);
410 s << t.format(TensorIOFormat::Plain());
The tensor base class.
Definition TensorForwardDeclarations.h:69
Namespace containing all symbols from the Eigen library.