Compute the minimum number of elements required in storage to hold a strided view described by dims, stride and offset.
| 21 | // Compute the minimum number of elements required in storage to hold |
| 22 | // a strided view described by dims, stride and offset. |
| 23 | static int64_t ComputeRequiredStorageSize(const std::vector<int64_t>& dims, |
| 24 | const std::vector<int64_t>& stride, |
| 25 | int64_t offset) { |
| 26 | int64_t required = offset; |
| 27 | for (size_t i = 0; i < dims.size(); ++i) { |
| 28 | if (dims[i] > 0) { |
| 29 | required += (dims[i] - 1) * stride[i]; |
| 30 | } |
| 31 | } |
| 32 | return required + 1; // +1 for the last element itself |
| 33 | } |
| 34 | |
| 35 | template <typename T, typename Context> |
| 36 | void SetKernel(const Context& dev_ctx, |