| 59 | }; |
| 60 | class SqueezeSizeComputer : public SizeComputer { |
| 61 | virtual bool onComputeSize(const MNN::Op* op, const std::vector<Tensor*>& inputs, |
| 62 | const std::vector<Tensor*>& outputs) const override { |
| 63 | MNN_ASSERT(1 == outputs.size()); |
| 64 | |
| 65 | const int* squeezeDim = nullptr; |
| 66 | int squeezeDimSize = 0; |
| 67 | if (nullptr != op->main_as_SqueezeParam()->squeezeDims()) { |
| 68 | squeezeDim = op->main_as_SqueezeParam()->squeezeDims()->data(); |
| 69 | squeezeDimSize = op->main_as_SqueezeParam()->squeezeDims()->size(); |
| 70 | } else if (inputs.size() > 1) { |
| 71 | squeezeDim = inputs[1]->host<int>(); |
| 72 | squeezeDimSize = inputs[1]->elementSize(); |
| 73 | } |
| 74 | uint32_t mask[MNN_MAX_TENSOR_DIM]; |
| 75 | ::memset(mask, 0, sizeof(mask)); |
| 76 | auto& ob = outputs[0]->buffer(); |
| 77 | auto& ib = inputs[0]->buffer(); |
| 78 | for (int i = 0; i < squeezeDimSize; i++) { |
| 79 | int axis = squeezeDim[i]; |
| 80 | if (axis < 0) { |
| 81 | axis += ib.dimensions; |
| 82 | } |
| 83 | if (axis < 0 || axis >= ib.dimensions) { |
| 84 | return false; |
| 85 | } |
| 86 | if (1 != ib.dim[axis].extent) { |
| 87 | MNN_ERROR("Cannot Squeeze dim[%d], 1 is expected, %d is got. input shape:", axis, ib.dim[axis].extent); |
| 88 | inputs[0]->printShape(); |
| 89 | return false; |
| 90 | } |
| 91 | mask[axis] = 1; |
| 92 | } |
| 93 | if (squeezeDimSize == 0) { |
| 94 | for (int i = 0; i < ib.dimensions; ++i) { |
| 95 | if (ib.dim[i].extent == 1) { |
| 96 | mask[i] = 1; |
| 97 | } |
| 98 | } |
| 99 | } |
| 100 | // Count actual unique squeezed dimensions from mask |
| 101 | int actualSqueeze = 0; |
| 102 | for (int i = 0; i < ib.dimensions; i++) { |
| 103 | if (mask[i]) { |
| 104 | actualSqueeze++; |
| 105 | } |
| 106 | } |
| 107 | |
| 108 | ob.dimensions = ib.dimensions - actualSqueeze; |
| 109 | int oDim = 0; |
| 110 | for (int i = 0; i < ib.dimensions; i++) { |
| 111 | if (mask[i] == 0) { |
| 112 | ob.dim[oDim].extent = ib.dim[i].extent; |
| 113 | oDim++; |
| 114 | } |
| 115 | } |
| 116 | ob.type = inputs[0]->buffer().type; |
| 117 | TensorUtils::getDescribe(outputs[0])->dimensionFormat = TensorUtils::getDescribe(inputs[0])->dimensionFormat; |
| 118 | return true; |
nothing calls this directly
no test coverage detected