| 25 | |
| 26 | template <typename T, typename Context> |
| 27 | void CrossGradKernel(const Context &dev_ctx, |
| 28 | const DenseTensor &x, |
| 29 | const DenseTensor &y, |
| 30 | const DenseTensor &out_grad, |
| 31 | int axis, |
| 32 | DenseTensor *x_grad, |
| 33 | DenseTensor *y_grad) { |
| 34 | auto &input_x = x; |
| 35 | auto &input_y = y; |
| 36 | auto &input_out_grad = out_grad; |
| 37 | auto *output_x_grad = x_grad; |
| 38 | auto *output_y_grad = y_grad; |
| 39 | int dim = axis; |
| 40 | auto input_x_dims = input_x.dims(); |
| 41 | if (dim != DDim::kMaxRank) { |
| 42 | PADDLE_ENFORCE_EQ( |
| 43 | dim < input_x_dims.size() && dim >= (0 - input_x_dims.size()), |
| 44 | true, |
| 45 | errors::OutOfRange( |
| 46 | "Attr(dim) is out of range, It's expected " |
| 47 | "to be in range of [-%d, %d]. But received Attr(dim) = %d.", |
| 48 | input_x_dims.size(), |
| 49 | input_x_dims.size() - 1, |
| 50 | dim)); |
| 51 | if (dim < 0) { |
| 52 | dim += input_x_dims.size(); |
| 53 | } |
| 54 | |
| 55 | PADDLE_ENFORCE_EQ( |
| 56 | input_x_dims[dim] == 3, |
| 57 | true, |
| 58 | errors::InvalidArgument( |
| 59 | "Input(X/Y).dims[dim] must be equal to 3. But received: " |
| 60 | "Input(X/Y).dims[dim] = [%d].", |
| 61 | input_x_dims[dim])); |
| 62 | } else { |
| 63 | for (auto i = 0; i < input_x_dims.size(); i++) { |
| 64 | if (input_x_dims[i] == 3) { |
| 65 | dim = i; |
| 66 | break; |
| 67 | } |
| 68 | } |
| 69 | PADDLE_ENFORCE_EQ( |
| 70 | dim == DDim::kMaxRank, |
| 71 | false, |
| 72 | errors::InvalidArgument("There must be at least one dimension 'd' " |
| 73 | "so that Input(X/Y).dims()[d] is equal to 3. " |
| 74 | "But received: Input(X/Y).dims() == [%s].", |
| 75 | input_x_dims)); |
| 76 | } |
| 77 | int64_t outer_loops = 1; |
| 78 | for (int i = 0; i < dim; i++) { |
| 79 | outer_loops *= input_x_dims[i]; |
| 80 | } |
| 81 | int64_t slice_size = 1; |
| 82 | for (int i = dim + 1; i < input_x_dims.size(); i++) { |
| 83 | slice_size *= input_x_dims[i]; |
| 84 | } |
nothing calls this directly
no test coverage detected