tensor.hpp
| Line | Branch | Exec | Source |
|---|---|---|---|
| 1 | #pragma once | ||
| 2 | |||
| 3 | #include <array> | ||
| 4 | |||
| 5 | namespace wala { | ||
| 6 | |||
| 7 | template <typename T, int NDIMS> struct tensor_view { | ||
| 8 | static_assert(NDIMS >= 0, "NDIMS must be nonnegative"); | ||
| 9 | |||
| 10 | protected: | ||
| 11 | std::array<int, NDIMS> shape; | ||
| 12 | std::array<int, NDIMS> strides; | ||
| 13 | T* data; | ||
| 14 | |||
| 15 | 54 | tensor_view(std::array<int, NDIMS> shape_, std::array<int, NDIMS> strides_, T* data_) : shape(shape_), strides(strides_), data(data_) {} | |
| 16 | |||
| 17 | public: | ||
| 18 | ✗ | tensor_view() : shape{0}, strides{0}, data(nullptr) {} | |
| 19 | |||
| 20 | protected: | ||
| 21 | ✗ | int flatten_index(std::array<int, NDIMS> idx) const { | |
| 22 | ✗ | int res = 0; | |
| 23 | ✗ | for (int i = 0; i < NDIMS; i++) { res += idx[i] * strides[i]; } | |
| 24 | ✗ | return res; | |
| 25 | } | ||
| 26 | 18 | int flatten_index_checked(std::array<int, NDIMS> idx) const { | |
| 27 | 18 | int res = 0; | |
| 28 |
2/2✓ Branch 30 → 3 taken 36 times.
✓ Branch 30 → 31 taken 18 times.
|
54 | for (int i = 0; i < NDIMS; i++) { |
| 29 | 36 | assert(0 <= idx[i] && idx[i] < shape[i]); | |
| 30 | 36 | res += idx[i] * strides[i]; | |
| 31 | } | ||
| 32 | 18 | return res; | |
| 33 | } | ||
| 34 | |||
| 35 | public: | ||
| 36 | 12 | T& operator[] (std::array<int, NDIMS> idx) const { | |
| 37 | #ifdef _GLIBCXX_DEBUG | ||
| 38 | 12 | return data[flatten_index_checked(idx)]; | |
| 39 | #else | ||
| 40 | return data[flatten_index(idx)]; | ||
| 41 | #endif | ||
| 42 | } | ||
| 43 | 6 | T& at(std::array<int, NDIMS> idx) const { | |
| 44 | 6 | return data[flatten_index_checked(idx)]; | |
| 45 | } | ||
| 46 | |||
| 47 | template <int D = NDIMS> | ||
| 48 | 24 | typename std::enable_if<(0 < D), tensor_view<T, NDIMS-1>>::type operator[] (int idx) const { | |
| 49 |
4/4std::enable_if<(0)<(1), wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> > const, 0> >::type wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> > const, 1>::operator[]<1>(int) const:
✓ Branch 15 → 16 taken 6 times.
std::enable_if<(0)<(2), wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> > const, 1> >::type wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> > const, 2>::operator[]<2>(int) const:
✓ Branch 17 → 18 taken 6 times.
std::enable_if<(0)<(1), wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> >, 0> >::type wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> >, 1>::operator[]<1>(int) const:
✓ Branch 15 → 16 taken 6 times.
std::enable_if<(0)<(2), wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> >, 1> >::type wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> >, 2>::operator[]<2>(int) const:
✓ Branch 17 → 18 taken 6 times.
|
84 | std::array<int, NDIMS-1> nshape; std::copy(shape.begin()+1, shape.end(), nshape.begin()); |
| 50 |
4/4std::enable_if<(0)<(1), wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> > const, 0> >::type wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> > const, 1>::operator[]<1>(int) const:
✓ Branch 33 → 34 taken 6 times.
std::enable_if<(0)<(2), wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> > const, 1> >::type wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> > const, 2>::operator[]<2>(int) const:
✓ Branch 37 → 38 taken 6 times.
std::enable_if<(0)<(1), wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> >, 0> >::type wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> >, 1>::operator[]<1>(int) const:
✓ Branch 33 → 34 taken 6 times.
std::enable_if<(0)<(2), wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> >, 1> >::type wala::tensor_view<std::__cxx11::basic_string<char, std::char_traits<char>, std::allocator<char> >, 2>::operator[]<2>(int) const:
✓ Branch 37 → 38 taken 6 times.
|
84 | std::array<int, NDIMS-1> nstrides; std::copy(strides.begin()+1, strides.end(), nstrides.begin()); |
| 51 | 48 | T* ndata = data + (strides[0] * idx); | |
| 52 | 60 | return tensor_view<T, NDIMS-1>(nshape, nstrides, ndata); | |
| 53 | } | ||
| 54 | template <int D = NDIMS> | ||
| 55 | ✗ | typename std::enable_if<(0 < D), tensor_view<T, NDIMS-1>>::type at(int idx) const { | |
| 56 | ✗ | assert(0 <= idx && idx < shape[0]); | |
| 57 | ✗ | return operator[](idx); | |
| 58 | } | ||
| 59 | |||
| 60 | template <int D = NDIMS> | ||
| 61 | 12 | typename std::enable_if<(0 == D), T&>::type operator * () const { | |
| 62 | 12 | return *data; | |
| 63 | } | ||
| 64 | |||
| 65 | template <typename U, int D> friend struct tensor_view; | ||
| 66 | template <typename U, int D> friend struct tensor; | ||
| 67 | }; | ||
| 68 | |||
| 69 | template <typename T, int NDIMS> struct tensor { | ||
| 70 | static_assert(NDIMS >= 0, "NDIMS must be nonnegative"); | ||
| 71 | |||
| 72 | protected: | ||
| 73 | std::array<int, NDIMS> shape; | ||
| 74 | std::array<int, NDIMS> strides; | ||
| 75 | int len; | ||
| 76 | T* data; | ||
| 77 | |||
| 78 | public: | ||
| 79 | ✗ | tensor() : shape{0}, strides{0}, len(0), data(nullptr) {} | |
| 80 | |||
| 81 | 1 | explicit tensor(std::array<int, NDIMS> shape_, const T& t = T()) { | |
| 82 | 1 | shape = shape_; | |
| 83 | 1 | len = 1; | |
| 84 |
2/2✓ Branch 29 → 8 taken 2 times.
✓ Branch 29 → 30 taken 1 time.
|
3 | for (int i = NDIMS-1; i >= 0; i--) { |
| 85 | 2 | strides[i] = len; | |
| 86 | 2 | len *= shape[i]; | |
| 87 | } | ||
| 88 |
3/4✓ Branch 33 → 34 taken 1 time.
✗ Branch 33 → 35 not taken.
✓ Branch 44 → 40 taken 6 times.
✓ Branch 44 → 45 taken 1 time.
|
7 | data = new T[len]; |
| 89 | 1 | std::fill(data, data + len, t); | |
| 90 | 1 | } | |
| 91 | |||
| 92 |
3/4✓ Branch 21 → 22 taken 2 times.
✗ Branch 21 → 23 not taken.
✓ Branch 32 → 28 taken 12 times.
✓ Branch 32 → 33 taken 2 times.
|
14 | tensor(const tensor& o) : shape(o.shape), strides(o.strides), len(o.len), data(new T[len]) { |
| 93 |
2/2✓ Branch 56 → 38 taken 12 times.
✓ Branch 56 → 57 taken 2 times.
|
14 | for (int i = 0; i < len; i++) { |
| 94 | 24 | data[i] = o.data[i]; | |
| 95 | } | ||
| 96 | 2 | } | |
| 97 | |||
| 98 | ✗ | tensor& operator=(tensor&& o) noexcept { | |
| 99 | using std::swap; | ||
| 100 | ✗ | swap(shape, o.shape); | |
| 101 | ✗ | swap(strides, o.strides); | |
| 102 | ✗ | swap(len, o.len); | |
| 103 | ✗ | swap(data, o.data); | |
| 104 | ✗ | return *this; | |
| 105 | } | ||
| 106 | ✗ | tensor(tensor&& o) : tensor() { | |
| 107 | ✗ | *this = std::move(o); | |
| 108 | } | ||
| 109 | ✗ | tensor& operator=(const tensor& o) { | |
| 110 | ✗ | return *this = tensor(o); | |
| 111 | } | ||
| 112 |
3/4✓ Branch 5 → 6 taken 3 times.
✗ Branch 5 → 40 not taken.
✓ Branch 20 → 21 taken 18 times.
✓ Branch 20 → 29 taken 3 times.
|
21 | ~tensor() { delete[] data; } |
| 113 | |||
| 114 | using view_t = tensor_view<T, NDIMS>; | ||
| 115 | 24 | view_t view() { | |
| 116 | 24 | return tensor_view<T, NDIMS>(shape, strides, data); | |
| 117 | } | ||
| 118 | ✗ | operator view_t() { | |
| 119 | ✗ | return view(); | |
| 120 | } | ||
| 121 | |||
| 122 | using const_view_t = tensor_view<const T, NDIMS>; | ||
| 123 | 6 | const_view_t view() const { | |
| 124 | 6 | return tensor_view<const T, NDIMS>(shape, strides, data); | |
| 125 | } | ||
| 126 | ✗ | operator const_view_t() const { | |
| 127 | ✗ | return view(); | |
| 128 | } | ||
| 129 | |||
| 130 | 12 | T& operator[] (std::array<int, NDIMS> idx) { return view()[idx]; } | |
| 131 | 6 | T& at(std::array<int, NDIMS> idx) { return view().at(idx); } | |
| 132 | ✗ | const T& operator[] (std::array<int, NDIMS> idx) const { return view()[idx]; } | |
| 133 | ✗ | const T& at(std::array<int, NDIMS> idx) const { return view().at(idx); } | |
| 134 | |||
| 135 | template <int D = NDIMS> | ||
| 136 | 6 | typename std::enable_if<(0 < D), tensor_view<T, NDIMS-1>>::type operator[] (int idx) { | |
| 137 |
1/1✓ Branch 6 → 7 taken 6 times.
|
6 | return view()[idx]; |
| 138 | } | ||
| 139 | template <int D = NDIMS> | ||
| 140 | ✗ | typename std::enable_if<(0 < D), tensor_view<T, NDIMS-1>>::type at(int idx) { | |
| 141 | ✗ | return view().at(idx); | |
| 142 | } | ||
| 143 | |||
| 144 | template <int D = NDIMS> | ||
| 145 | 6 | typename std::enable_if<(0 < D), tensor_view<const T, NDIMS-1>>::type operator[] (int idx) const { | |
| 146 |
1/1✓ Branch 6 → 7 taken 6 times.
|
6 | return view()[idx]; |
| 147 | } | ||
| 148 | template <int D = NDIMS> | ||
| 149 | ✗ | typename std::enable_if<(0 < D), tensor_view<const T, NDIMS-1>>::type at(int idx) const { | |
| 150 | ✗ | return view().at(idx); | |
| 151 | } | ||
| 152 | |||
| 153 | template <int D = NDIMS> | ||
| 154 | ✗ | typename std::enable_if<(0 == D), T&>::type operator * () { | |
| 155 | ✗ | return *view(); | |
| 156 | } | ||
| 157 | template <int D = NDIMS> | ||
| 158 | ✗ | typename std::enable_if<(0 == D), const T&>::type operator * () const { | |
| 159 | ✗ | return *view(); | |
| 160 | } | ||
| 161 | }; | ||
| 162 | |||
| 163 | } // namespace wala | ||
| 164 |