// This file is part of Eigen, a lightweight C++ template library // for linear algebra. // // Copyright (C) 2014 Benoit Steiner // // This Source Code Form is subject to the terms of the Mozilla // Public License v. 2.0. If a copy of the MPL was not distributed // with this file, You can obtain one at http://mozilla.org/MPL/2.0/. #ifndef EIGEN_CXX11_TENSOR_TENSOR_IO_H #define EIGEN_CXX11_TENSOR_TENSOR_IO_H #include "./InternalHeaderCheck.h" namespace Eigen { struct TensorIOFormat; namespace internal { template struct TensorPrinter; } struct TensorIOFormat { TensorIOFormat(const std::vector& _separator, const std::vector& _prefix, const std::vector& _suffix, int _precision = StreamPrecision, int _flags = 0, const std::string& _tenPrefix = "", const std::string& _tenSuffix = "", const char _fill = ' ') : tenPrefix(_tenPrefix), tenSuffix(_tenSuffix), prefix(_prefix), suffix(_suffix), separator(_separator), fill(_fill), precision(_precision), flags(_flags) { init_spacer(); } TensorIOFormat(int _precision = StreamPrecision, int _flags = 0, const std::string& _tenPrefix = "", const std::string& _tenSuffix = "", const char _fill = ' ') : tenPrefix(_tenPrefix), tenSuffix(_tenSuffix), fill(_fill), precision(_precision), flags(_flags) { // default values of prefix, suffix and separator prefix = {"", "["}; suffix = {"", "]"}; separator = {", ", "\n"}; init_spacer(); } void init_spacer() { if ((flags & DontAlignCols)) return; spacer.resize(prefix.size()); spacer[0] = ""; int i = int(tenPrefix.length()) - 1; while (i >= 0 && tenPrefix[i] != '\n') { spacer[0] += ' '; i--; } for (std::size_t k = 1; k < prefix.size(); k++) { int i = int(prefix[k].length()) - 1; while (i >= 0 && prefix[k][i] != '\n') { spacer[k] += ' '; i--; } } } static inline const TensorIOFormat Numpy() { std::vector prefix = {"", "["}; std::vector suffix = {"", "]"}; std::vector separator = {" ", "\n"}; return TensorIOFormat(separator, prefix, suffix, StreamPrecision, 0, "[", "]"); } static inline const TensorIOFormat Plain() { std::vector separator = {" ", "\n", "\n", ""}; std::vector prefix = {""}; std::vector suffix = {""}; return TensorIOFormat(separator, prefix, suffix, StreamPrecision, 0, "", "", ' '); } static inline const TensorIOFormat Native() { std::vector separator = {", ", ",\n", "\n"}; std::vector prefix = {"", "{"}; std::vector suffix = {"", "}"}; return TensorIOFormat(separator, prefix, suffix, StreamPrecision, 0, "{", "}", ' '); } static inline const TensorIOFormat Legacy() { TensorIOFormat LegacyFormat(StreamPrecision, 0, "", "", ' '); LegacyFormat.legacy_bit = true; return LegacyFormat; } std::string tenPrefix; std::string tenSuffix; std::vector prefix; std::vector suffix; std::vector separator; char fill; int precision; int flags; std::vector spacer{}; bool legacy_bit = false; }; template class TensorWithFormat; // specialize for Layout=ColMajor, Layout=RowMajor and rank=0. template class TensorWithFormat { public: TensorWithFormat(const T& tensor, const TensorIOFormat& format) : t_tensor(tensor), t_format(format) {} friend std::ostream& operator<<(std::ostream& os, const TensorWithFormat& wf) { // Evaluate the expression if needed typedef TensorEvaluator, DefaultDevice> Evaluator; TensorForcedEvalOp eval = wf.t_tensor.eval(); Evaluator tensor(eval, DefaultDevice()); tensor.evalSubExprsIfNeeded(NULL); internal::TensorPrinter::run(os, tensor, wf.t_format); // Cleanup. tensor.cleanup(); return os; } protected: T t_tensor; TensorIOFormat t_format; }; template class TensorWithFormat { public: TensorWithFormat(const T& tensor, const TensorIOFormat& format) : t_tensor(tensor), t_format(format) {} friend std::ostream& operator<<(std::ostream& os, const TensorWithFormat& wf) { // Switch to RowMajor storage and print afterwards typedef typename T::Index Index; std::array shuffle; std::array id; std::iota(id.begin(), id.end(), Index(0)); std::copy(id.begin(), id.end(), shuffle.rbegin()); auto tensor_row_major = wf.t_tensor.swap_layout().shuffle(shuffle); // Evaluate the expression if needed typedef TensorEvaluator, DefaultDevice> Evaluator; TensorForcedEvalOp eval = tensor_row_major.eval(); Evaluator tensor(eval, DefaultDevice()); tensor.evalSubExprsIfNeeded(NULL); internal::TensorPrinter::run(os, tensor, wf.t_format); // Cleanup. tensor.cleanup(); return os; } protected: T t_tensor; TensorIOFormat t_format; }; template class TensorWithFormat { public: TensorWithFormat(const T& tensor, const TensorIOFormat& format) : t_tensor(tensor), t_format(format) {} friend std::ostream& operator<<(std::ostream& os, const TensorWithFormat& wf) { // Evaluate the expression if needed typedef TensorEvaluator, DefaultDevice> Evaluator; TensorForcedEvalOp eval = wf.t_tensor.eval(); Evaluator tensor(eval, DefaultDevice()); tensor.evalSubExprsIfNeeded(NULL); internal::TensorPrinter::run(os, tensor, wf.t_format); // Cleanup. tensor.cleanup(); return os; } protected: T t_tensor; TensorIOFormat t_format; }; namespace internal { template struct TensorPrinter { static void run(std::ostream& s, const Tensor& _t, const TensorIOFormat& fmt) { typedef typename internal::remove_const::type Scalar; typedef typename Tensor::Index Index; static const int layout = Tensor::Layout; // backwards compatibility case: print tensor after reshaping to matrix of size dim(0) x // (dim(1)*dim(2)*...*dim(rank-1)). if (fmt.legacy_bit) { const Index total_size = internal::array_prod(_t.dimensions()); if (total_size > 0) { const Index first_dim = Eigen::internal::array_get<0>(_t.dimensions()); Map > matrix(_t.data(), first_dim, total_size / first_dim); s << matrix; return; } } assert(layout == RowMajor); typedef typename conditional::value || is_same::value || is_same::value || is_same::value, int, typename conditional >::value || is_same >::value || is_same >::value || is_same >::value, std::complex, const Scalar&>::type>::type PrintType; const Index total_size = array_prod(_t.dimensions()); std::streamsize explicit_precision; if (fmt.precision == StreamPrecision) { explicit_precision = 0; } else if (fmt.precision == FullPrecision) { if (NumTraits::IsInteger) { explicit_precision = 0; } else { explicit_precision = significant_decimals_impl::run(); } } else { explicit_precision = fmt.precision; } std::streamsize old_precision = 0; if (explicit_precision) old_precision = s.precision(explicit_precision); Index width = 0; bool align_cols = !(fmt.flags & DontAlignCols); if (align_cols) { // compute the largest width for (Index i = 0; i < total_size; i++) { std::stringstream sstr; sstr.copyfmt(s); sstr << static_cast(_t.data()[i]); width = std::max(width, Index(sstr.str().length())); } } std::streamsize old_width = s.width(); char old_fill_character = s.fill(); s << fmt.tenPrefix; for (Index i = 0; i < total_size; i++) { std::array is_at_end{}; std::array is_at_begin{}; // is the ith element the end of an coeff (always true), of a row, of a matrix, ...? for (std::size_t k = 0; k < rank; k++) { if ((i + 1) % (std::accumulate(_t.dimensions().rbegin(), _t.dimensions().rbegin() + k, 1, std::multiplies())) == 0) { is_at_end[k] = true; } } // is the ith element the begin of an coeff (always true), of a row, of a matrix, ...? for (std::size_t k = 0; k < rank; k++) { if (i % (std::accumulate(_t.dimensions().rbegin(), _t.dimensions().rbegin() + k, 1, std::multiplies())) == 0) { is_at_begin[k] = true; } } // do we have a line break? bool is_at_begin_after_newline = false; for (std::size_t k = 0; k < rank; k++) { if (is_at_begin[k]) { std::size_t separator_index = (k < fmt.separator.size()) ? k : fmt.separator.size() - 1; if (fmt.separator[separator_index].find('\n') != std::string::npos) { is_at_begin_after_newline = true; } } } bool is_at_end_before_newline = false; for (std::size_t k = 0; k < rank; k++) { if (is_at_end[k]) { std::size_t separator_index = (k < fmt.separator.size()) ? k : fmt.separator.size() - 1; if (fmt.separator[separator_index].find('\n') != std::string::npos) { is_at_end_before_newline = true; } } } std::stringstream suffix, prefix, separator; for (std::size_t k = 0; k < rank; k++) { std::size_t suffix_index = (k < fmt.suffix.size()) ? k : fmt.suffix.size() - 1; if (is_at_end[k]) { suffix << fmt.suffix[suffix_index]; } } for (std::size_t k = 0; k < rank; k++) { std::size_t separator_index = (k < fmt.separator.size()) ? k : fmt.separator.size() - 1; if (is_at_end[k] and (!is_at_end_before_newline or fmt.separator[separator_index].find('\n') != std::string::npos)) { separator << fmt.separator[separator_index]; } } for (std::size_t k = 0; k < rank; k++) { std::size_t spacer_index = (k < fmt.spacer.size()) ? k : fmt.spacer.size() - 1; if (i != 0 and is_at_begin_after_newline and (!is_at_begin[k] or k == 0)) { prefix << fmt.spacer[spacer_index]; } } for (int k = rank - 1; k >= 0; k--) { std::size_t prefix_index = (static_cast(k) < fmt.prefix.size()) ? k : fmt.prefix.size() - 1; if (is_at_begin[k]) { prefix << fmt.prefix[prefix_index]; } } s << prefix.str(); if (width) { s.fill(fmt.fill); s.width(width); s << std::right; } s << _t.data()[i]; s << suffix.str(); if (i < total_size - 1) { s << separator.str(); } } s << fmt.tenSuffix; if (explicit_precision) s.precision(old_precision); if (width) { s.fill(old_fill_character); s.width(old_width); } } }; template struct TensorPrinter { static void run(std::ostream& s, const Tensor& _t, const TensorIOFormat& fmt) { typedef typename Tensor::Scalar Scalar; std::streamsize explicit_precision; if (fmt.precision == StreamPrecision) { explicit_precision = 0; } else if (fmt.precision == FullPrecision) { if (NumTraits::IsInteger) { explicit_precision = 0; } else { explicit_precision = significant_decimals_impl::run(); } } else { explicit_precision = fmt.precision; } std::streamsize old_precision = 0; if (explicit_precision) old_precision = s.precision(explicit_precision); s << fmt.tenPrefix << _t.coeff(0) << fmt.tenSuffix; if (explicit_precision) s.precision(old_precision); } }; } // end namespace internal template std::ostream& operator<<(std::ostream& s, const TensorBase& t) { s << t.format(TensorIOFormat::Plain()); return s; } } // end namespace Eigen #endif // EIGEN_CXX11_TENSOR_TENSOR_IO_H