| 27 | namespace phi { |
| 28 | |
| 29 | std::pair<std::vector<int64_t>, std::vector<int64_t>> GetReduceDims( |
| 30 | const DDim& src_dim, const DDim& dst_dim) { |
| 31 | std::vector<int64_t> reduce_dims, new_dims; |
| 32 | auto pre_dims = src_dim.size() - dst_dim.size(); |
| 33 | for (auto i = 0; i < pre_dims; ++i) { |
| 34 | reduce_dims.push_back(i); |
| 35 | } |
| 36 | |
| 37 | for (auto i = pre_dims; i < src_dim.size(); ++i) { |
| 38 | if (dst_dim[i - pre_dims] == 1 && src_dim[i] != 1) { |
| 39 | reduce_dims.push_back(i); |
| 40 | } else { |
| 41 | new_dims.push_back(dst_dim[i - pre_dims]); |
| 42 | } |
| 43 | } |
| 44 | return {reduce_dims, new_dims}; |
| 45 | } |
| 46 | |
| 47 | template <typename T, typename Context> |
| 48 | void DistGradKernel(const Context& dev_ctx, |
no test coverage detected