| 108 | } |
| 109 | |
| 110 | static shape mask_shape(const shape& s, const std::vector<std::size_t>& lens) |
| 111 | { |
| 112 | assert(s.lens().size() == lens.size()); |
| 113 | std::vector<std::size_t> rstrides(lens.size()); |
| 114 | std::size_t stride = 1; |
| 115 | for(std::size_t i = lens.size() - 1; i < lens.size(); i--) |
| 116 | { |
| 117 | if(lens[i] == s.lens()[i]) |
| 118 | { |
| 119 | rstrides[i] = stride; |
| 120 | stride *= lens[i]; |
| 121 | } |
| 122 | else if(lens[i] != 1 and s.lens()[i] != 1) |
| 123 | { |
| 124 | return shape{}; |
| 125 | } |
| 126 | } |
| 127 | return shape{s.type(), lens, rstrides}; |
| 128 | } |
| 129 | |
| 130 | std::vector<shape> reduce_dims(const std::vector<shape>& shapes) |
| 131 | { |
no test coverage detected