| 11 | #include "core/OpCommonUtils.hpp" |
| 12 | namespace MNN { |
| 13 | static std::vector<std::tuple<int, int, int>> _computeReduceDims(const std::vector<Tensor*>& inputs, |
| 14 | std::vector<int>& axises) { |
| 15 | |
| 16 | auto totalSize = TensorUtils::getRawSize(inputs[0]); |
| 17 | if (axises.empty()) { |
| 18 | return {std::make_tuple(1, totalSize, 1)}; |
| 19 | } |
| 20 | for (int i = 0; i < axises.size(); ++i) { |
| 21 | if (axises[i] < 0) { |
| 22 | if (axises[i] < 0) { |
| 23 | return {std::make_tuple(1, totalSize, 1)}; |
| 24 | } |
| 25 | } |
| 26 | } |
| 27 | // Cache for input's dims |
| 28 | std::vector<int> lengths(inputs[0]->dimensions()); |
| 29 | for (int i = 0; i < lengths.size(); ++i) { |
| 30 | lengths[i] = inputs[0]->length(i); |
| 31 | } |
| 32 | std::vector<std::pair<int, int>> groupAxises; |
| 33 | { |
| 34 | // Merge adj axis |
| 35 | std::sort(axises.begin(), axises.end()); |
| 36 | int lastAxis = axises[0]; |
| 37 | int length = 1; |
| 38 | int start = axises[0]; |
| 39 | for (int i = 1; i < axises.size(); ++i) { |
| 40 | // MNN_PRINT("%d - %d\n", axises[i], lastAxis); |
| 41 | if (axises[i] - lastAxis == 1) { |
| 42 | length++; |
| 43 | } else { |
| 44 | groupAxises.emplace_back(std::make_pair(start, length)); |
| 45 | length = 1; |
| 46 | start = axises[i]; |
| 47 | } |
| 48 | lastAxis = axises[i]; |
| 49 | } |
| 50 | groupAxises.emplace_back(std::make_pair(start, length)); |
| 51 | } |
| 52 | |
| 53 | // Compute inside-outside-axis |
| 54 | std::vector<std::tuple<int, int, int>> result; |
| 55 | |
| 56 | for (int i = 0; i < groupAxises.size(); ++i) { |
| 57 | int outsideSize = 1; |
| 58 | int insideSize = 1; |
| 59 | int axisSize = 1; |
| 60 | auto start = groupAxises[i].first; |
| 61 | auto length = groupAxises[i].second; |
| 62 | if (start >= (int)lengths.size()) { |
| 63 | break; |
| 64 | } |
| 65 | for (int j = 0; j < start; ++j) { |
| 66 | outsideSize *= lengths[j]; |
| 67 | } |
| 68 | for (int j = start; j < start + length; ++j) { |
| 69 | if (j >= (int)lengths.size()) { |
| 70 | break; |