| 283 | return kernel_code |
| 284 | |
| 285 | def gen_top_api(kernel_shapes): |
| 286 | |
| 287 | kernel_code = "void ggml_preprocessor(int m, int k, void* B, void* LUT_Scales, void* QLUT) {{\n\ |
| 288 | if (m == {0} && k == {1}) {{\n\ |
| 289 | preprocessor_k<{1}>(B, LUT_Scales, QLUT);\n\ |
| 290 | }}\n\ |
| 291 | ".format(kernel_shapes[0][0], kernel_shapes[0][1]) |
| 292 | for i in range(1, len(kernel_shapes)): |
| 293 | kernel_code = "".join([kernel_code, " else if (m == {0} && k == {1}) {{\n\ |
| 294 | preprocessor_k<{1}>(B, LUT_Scales, QLUT);\n\ |
| 295 | }}\n".format(kernel_shapes[i][0], kernel_shapes[i][1])]) |
| 296 | kernel_code = "".join([kernel_code, "}\n"]) |
| 297 | kernel_code = "".join([kernel_code, "void ggml_qgemm_lut(int m, int k, void* A, void* LUT, void* Scales, void* LUT_Scales, void* C) {{\n\ |
| 298 | if (m == {0} && k == {1}) {{\n\ |
| 299 | qgemm_lut_{0}_{1}(A, LUT, Scales, LUT_Scales, C);\n\ |
| 300 | }}\n\ |
| 301 | ".format(kernel_shapes[0][0], kernel_shapes[0][1])]) |
| 302 | for i in range(1, len(kernel_shapes)): |
| 303 | kernel_code = "".join([kernel_code, " else if (m == {0} && k == {1}) {{\n\ |
| 304 | qgemm_lut_{0}_{1}(A, LUT, Scales, LUT_Scales, C);\n\ |
| 305 | }}\n\ |
| 306 | ".format(kernel_shapes[i][0], kernel_shapes[i][1])]) |
| 307 | kernel_code = "".join([kernel_code, "}\n"]) |
| 308 | return kernel_code |
| 309 | |
| 310 | def gen_preprocess_code(): |
| 311 | kernel_code = "\n\ |