| 27 | |
| 28 | template <typename Context> |
| 29 | void SliceStridedKernel(const Context& dev_ctx, |
| 30 | const DenseTensor& input, |
| 31 | const std::vector<int64_t>& axes, |
| 32 | const IntArray& starts_arr, |
| 33 | const IntArray& ends_arr, |
| 34 | const std::vector<int64_t>& infer_flags, |
| 35 | const std::vector<int64_t>& decrease_axis, |
| 36 | DenseTensor* out) { |
| 37 | if (!FLAGS_use_stride_kernel) { |
| 38 | PADDLE_THROW(common::errors::Fatal( |
| 39 | "FLAGS_use_stride_kernel is closed. Strided kernel " |
| 40 | "be called, something wrong has happened!")); |
| 41 | } |
| 42 | std::vector<int64_t> starts = starts_arr.GetData(); |
| 43 | std::vector<int64_t> ends = ends_arr.GetData(); |
| 44 | const auto& in_dims = input.dims(); |
| 45 | |
| 46 | auto new_axes = axes; |
| 47 | for (auto& item : new_axes) { |
| 48 | if (item < 0) { |
| 49 | item = std::max(int64_t(0), item + int64_t(in_dims.size())); |
| 50 | } |
| 51 | } |
| 52 | // axis = 0, dim_value = 3, st[0]=0, ed[0]=4 |
| 53 | // The step seems to be regarded as 1 here |
| 54 | funcs::CheckAndUpdateSliceAttrs<int64_t>( |
| 55 | in_dims, new_axes, &starts, &ends, nullptr, nullptr); |
| 56 | |
| 57 | std::vector<int64_t> output_dims = vectorize<int64_t>(input.dims()); |
| 58 | std::vector<int64_t> output_stride = vectorize<int64_t>(input.strides()); |
| 59 | int64_t output_offset = static_cast<int64_t>(input.offset()); |
| 60 | |
| 61 | for (size_t i = 0; i < new_axes.size(); ++i) { |
| 62 | output_offset = static_cast<int64_t>( |
| 63 | output_offset + |
| 64 | starts[i] * output_stride[new_axes[i]] * SizeOf(out->dtype())); |
| 65 | output_dims[new_axes[i]] = std::abs(ends[i] - starts[i]); |
| 66 | } |
| 67 | |
| 68 | std::vector<uint8_t> decrease_flag(output_dims.size(), 0); |
| 69 | if (!decrease_axis.empty()) { |
| 70 | for (auto axis : decrease_axis) { |
| 71 | decrease_flag[axis] = 1; |
| 72 | } |
| 73 | |
| 74 | std::vector<int64_t> new_shape; |
| 75 | std::vector<int64_t> new_stride; |
| 76 | for (size_t i = 0; i < output_dims.size(); ++i) { |
| 77 | if (decrease_flag[i] == 0) { |
| 78 | new_shape.push_back(output_dims[i]); |
| 79 | new_stride.push_back(output_stride[i]); |
| 80 | } |
| 81 | } |
| 82 | output_dims = new_shape; |
| 83 | output_stride = new_stride; |
| 84 | } |
| 85 | |
| 86 | auto meta = out->meta(); |
nothing calls this directly
no test coverage detected