| 23 | |
| 24 | template <typename Context> |
| 25 | void TransposeStridedKernel(const Context& dev_ctx, |
| 26 | const DenseTensor& x, |
| 27 | const std::vector<int>& axis, |
| 28 | DenseTensor* out) { |
| 29 | if (!FLAGS_use_stride_kernel) { |
| 30 | PADDLE_THROW(common::errors::Fatal( |
| 31 | "FLAGS_use_stride_kernel is closed. Strided kernel " |
| 32 | "be called, something wrong has happened!")); |
| 33 | } |
| 34 | size_t x_rank = x.dims().size(); |
| 35 | std::vector<int> formatted_axis = axis; |
| 36 | for (size_t i = 0; i < axis.size(); i++) { |
| 37 | if (axis[i] < 0) { |
| 38 | formatted_axis[i] = static_cast<int>(axis[i] + x_rank); |
| 39 | } |
| 40 | } |
| 41 | |
| 42 | auto meta = out->meta(); |
| 43 | auto in_stride = x.strides(); |
| 44 | meta.strides = in_stride; |
| 45 | for (int i = 0; i < static_cast<int>(formatted_axis.size()); i++) { |
| 46 | meta.strides[i] = in_stride[formatted_axis[i]]; |
| 47 | } |
| 48 | meta.offset = x.offset(); |
| 49 | |
| 50 | out->set_meta(meta); |
| 51 | out->ResetHolder(x.Holder()); |
| 52 | out->ShareInplaceVersionCounterWith(x); |
| 53 | } |
| 54 | |
| 55 | } // namespace phi |
| 56 | |