MCPcopy Create free account
hub / github.com/PaddlePaddle/Paddle / SliceStridedKernel

Function SliceStridedKernel

paddle/phi/kernels/stride/slice_kernel.cc:29–104  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

27
28template <typename Context>
29void 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();

Callers

nothing calls this directly

Calls 15

SizeOfFunction · 0.85
metaMethod · 0.80
ResetHolderMethod · 0.80
HolderMethod · 0.80
maxFunction · 0.50
absFunction · 0.50
DDimClass · 0.50
GetDataMethod · 0.45
dimsMethod · 0.45
sizeMethod · 0.45
stridesMethod · 0.45
offsetMethod · 0.45

Tested by

no test coverage detected