| 236 | } |
| 237 | |
| 238 | shape make_bcast_shape(const shape& input_shape, const std::vector<std::size_t>& bcast_lens) |
| 239 | { |
| 240 | assert(not input_shape.dynamic()); |
| 241 | auto offset = bcast_lens.size() - input_shape.ndim(); |
| 242 | std::vector<size_t> bcast_strides(bcast_lens.size(), 0); |
| 243 | for(std::ptrdiff_t i : reverse(range(input_shape.ndim()))) |
| 244 | { |
| 245 | if(bcast_lens.at(i + offset) == input_shape.lens()[i]) |
| 246 | { |
| 247 | bcast_strides.at(i + offset) = input_shape.strides()[i]; |
| 248 | } |
| 249 | } |
| 250 | return shape{input_shape.type(), bcast_lens, bcast_strides}; |
| 251 | } |
| 252 | |
| 253 | } // namespace MIGRAPHX_INLINE_NS |
| 254 | } // namespace migraphx |