MCPcopy Create free account
hub / github.com/MegEngine/MegCC / GetKernelBody

Method GetKernelBody

compiler/lib/KernelGen/BareMetal/MatrixInv.cpp:42–118  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

40}
41
42std::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 }

Callers

nothing calls this directly

Calls 1

GenCommonRetFunction · 0.85

Tested by

no test coverage detected