| 191 | } |
| 192 | |
| 193 | void Matrix::invert(Tensor* dst, const Tensor* src) { |
| 194 | MNN_ASSERT(2 == src->buffer().dimensions); |
| 195 | const int N0 = src->buffer().dim[0].extent; |
| 196 | MNN_ASSERT(N0 == src->buffer().dim[1].extent); |
| 197 | |
| 198 | int i, j, k; |
| 199 | float max, temp; |
| 200 | std::shared_ptr<Tensor> tempMat(Matrix::create(N0, N0)); |
| 201 | ::memcpy(tempMat->buffer().host, src->buffer().host, src->size()); |
| 202 | const auto tempData = tempMat->host<float>(); |
| 203 | const auto dstData = dst->host<float>(); |
| 204 | for (i = 0; i < N0; ++i) { |
| 205 | for (j = 0; j < N0; ++j) { |
| 206 | *(dstData + i * N0 + j) = (i == j) ? 1.0f : 0.0f; |
| 207 | } |
| 208 | } |
| 209 | |
| 210 | for (i = 0; i < N0; ++i) { |
| 211 | max = *(tempData + i * N0 + i); |
| 212 | k = i; |
| 213 | for (j = i + 1; j < N0; ++j) { |
| 214 | auto data1 = *(tempData + j * N0 + i); |
| 215 | if (fabs(data1) > fabs(max)) { |
| 216 | max = data1; |
| 217 | k = j; |
| 218 | } |
| 219 | } |
| 220 | if (k != i) { |
| 221 | for (j = 0; j < N0; ++j) { |
| 222 | temp = *(tempData + i * N0 + j); |
| 223 | *(tempData + i * N0 + j) = *(tempData + k * N0 + j); |
| 224 | *(tempData + k * N0 + j) = temp; |
| 225 | temp = *(dstData + i * N0 + j); |
| 226 | *(dstData + i * N0 + j) = *(dstData + k * N0 + j); |
| 227 | *(dstData + k * N0 + j) = temp; |
| 228 | } |
| 229 | } |
| 230 | if (*(tempData + i * N0 + i) == 0) { |
| 231 | MNN_PRINT("This matrix have no inverse!\n"); |
| 232 | return; |
| 233 | } |
| 234 | temp = *(tempData + i * N0 + i); |
| 235 | |
| 236 | for (j = 0; j < N0; ++j) { |
| 237 | *(tempData + i * N0 + j) = *(tempData + i * N0 + j) / temp; |
| 238 | *(dstData + i * N0 + j) = *(dstData + i * N0 + j) / temp; |
| 239 | } |
| 240 | |
| 241 | for (j = 0; j < N0; ++j) { |
| 242 | if (j != i) { |
| 243 | temp = *(tempData + j * N0 + i); |
| 244 | for (k = 0; k < N0; ++k) { |
| 245 | *(tempData + j * N0 + k) = *(tempData + j * N0 + k) - *(tempData + i * N0 + k) * temp; |
| 246 | *(dstData + j * N0 + k) = *(dstData + j * N0 + k) - *(dstData + i * N0 + k) * temp; |
| 247 | } |
| 248 | } |
| 249 | } |
| 250 | } |