| 40 | } |
| 41 | |
| 42 | std::string MatrixInvKernel::GetKernelBody(TContext* context) const { |
| 43 | std::stringstream ss; |
| 44 | ss << R"( |
| 45 | #include <math.h> |
| 46 | #include <string.h> |
| 47 | )"; |
| 48 | ss << GenCommonRet() << " " << GetKernelSignature(context); |
| 49 | std::string body_temp = R"({ |
| 50 | float* a_data = (float*)inputs[0]->ptr; |
| 51 | float* c_data = (float*)outputs[0]->ptr; |
| 52 | TINYNN_ASSERT(a_data); |
| 53 | TINYNN_ASSERT(c_data); |
| 54 | const Tensor* a_tensor = inputs[0]; |
| 55 | const Tensor* c_tensor = outputs[0]; |
| 56 | const Layout a_layout = a_tensor->layout; |
| 57 | const int n = a_layout.dims[a_layout.nr_dim - 1]; |
| 58 | const int ld_buffer = 2 * n; |
| 59 | int batch = 1; |
| 60 | for (int i = 0; i < a_layout.nr_dim - 2; ++i) { |
| 61 | batch *= a_layout.dims[i]; |
| 62 | } |
| 63 | float* src_buffer = (float*)(workspace->ptr); |
| 64 | for (int b_idx = 0; b_idx < batch; ++b_idx){ |
| 65 | float* batch_src = a_data + b_idx * n * n; |
| 66 | float* batch_dst = c_data + b_idx * n * n; |
| 67 | for (int row = 0; row < n; ++row){ |
| 68 | memcpy(&src_buffer[row * ld_buffer], batch_src + row * n, sizeof(float) * n); |
| 69 | memset(&src_buffer[row * ld_buffer] + n, 0, sizeof(float) * n); |
| 70 | src_buffer[row * ld_buffer + n + row] = 1; |
| 71 | } |
| 72 | |
| 73 | for (int out_row = 0; out_row < n; ++out_row){ |
| 74 | float abs_max = 0.f; |
| 75 | int select_row = out_row; |
| 76 | for (int row = out_row; row < n; ++row){ |
| 77 | float abs_val = fabsf(src_buffer[row * ld_buffer + out_row]); |
| 78 | if (abs_val > abs_max){ |
| 79 | abs_max = abs_val; |
| 80 | select_row = row; |
| 81 | } |
| 82 | } |
| 83 | TINYNN_ASSERT(abs_max > 1e-7); |
| 84 | for(int col = 0; col < 2 * n; ++col){ |
| 85 | float temp = src_buffer[out_row * ld_buffer + col]; |
| 86 | src_buffer[out_row * ld_buffer + col] = src_buffer[select_row * ld_buffer + col]; |
| 87 | src_buffer[select_row * ld_buffer + col] = temp; |
| 88 | } |
| 89 | |
| 90 | // substract pivot row from other rows |
| 91 | float* pivot_row_ptr = &src_buffer[out_row * ld_buffer]; |
| 92 | for (int row = 0; row < n; ++row) { |
| 93 | if (row == out_row) { |
| 94 | continue; |
| 95 | } |
| 96 | float inv_pivot = -src_buffer[row * ld_buffer + out_row] / pivot_row_ptr[out_row]; |
| 97 | for (int col = out_row; col < n * 2; ++col) { |
| 98 | src_buffer[row * ld_buffer + col] += pivot_row_ptr[col] * inv_pivot; |
| 99 | } |
nothing calls this directly
no test coverage detected