| 624 | return kernel_code |
| 625 | |
| 626 | def gen_transform_code(kernel_shapes): |
| 627 | kernel_code = "\n\ |
| 628 | void ggml_bitnet_transform_tensor(struct ggml_tensor * tensor) {\n\ |
| 629 | if (!(is_type_supported(tensor->type) && tensor->backend == GGML_BACKEND_TYPE_CPU && tensor->extra == nullptr)) {\n\ |
| 630 | return;\n\ |
| 631 | }\n\ |
| 632 | \n\ |
| 633 | int k = tensor->ne[0];\n\ |
| 634 | int m = tensor->ne[1];\n\ |
| 635 | const int lut_scales_size = 1;\n\ |
| 636 | int bk = 0;\n\ |
| 637 | int bm = 0;\n" |
| 638 | |
| 639 | kernel_code = "".join([kernel_code, "\n\ |
| 640 | if (m == {0} && k == {1}) {{\n\ |
| 641 | bm = BM{0}_{1};\n\ |
| 642 | bk = BBK{0}_{1};\n\ |
| 643 | }}\n".format(kernel_shapes[0][0], kernel_shapes[0][1])]) |
| 644 | |
| 645 | for i in range(1, len(kernel_shapes)): |
| 646 | kernel_code = "".join([kernel_code, "else if (m == {0} && k == {1}) {{\n\ |
| 647 | bm = BM{0}_{1};\n\ |
| 648 | bk = BBK{0}_{1};\n\ |
| 649 | }}\n".format(kernel_shapes[i][0], kernel_shapes[i][1])]) |
| 650 | |
| 651 | kernel_code = "".join([kernel_code, "\n\ |
| 652 | const int n_tile_num = m / bm;\n\ |
| 653 | const int BK = bk;\n\ |
| 654 | uint8_t * qweights;\n\ |
| 655 | bitnet_float_type * scales;\n\ |
| 656 | \n\ |
| 657 | scales = (bitnet_float_type *) aligned_malloc(sizeof(bitnet_float_type));\n\ |
| 658 | qweights = (uint8_t *) tensor->data;\n\ |
| 659 | int nbytes = (k - 256) * m / 3 * 5 / 8 + 256 * m / 2 * 4 / 8;\n\ |
| 660 | if (nbytes % 32 != 0) nbytes = 32 - nbytes % 32 + nbytes;\n\ |
| 661 | float * i2_scales = (float * )(qweights + nbytes);\n\ |
| 662 | scales[0] = (bitnet_float_type) i2_scales[0];\n\ |
| 663 | \n\ |
| 664 | tensor->extra = bitnet_tensor_extras + bitnet_tensor_extras_index;\n\ |
| 665 | bitnet_tensor_extras[bitnet_tensor_extras_index++] = {\n\ |
| 666 | /* .lut_scales_size = */ lut_scales_size,\n\ |
| 667 | /* .BK = */ BK,\n\ |
| 668 | /* .n_tile_num = */ n_tile_num,\n\ |
| 669 | /* .qweights = */ qweights,\n\ |
| 670 | /* .scales = */ scales\n\ |
| 671 | };\n\ |
| 672 | }\n"]) |
| 673 | |
| 674 | return kernel_code |
| 675 | |
| 676 | def get_three_k_two_k(K, bk): |
| 677 | bk_num = K // bk |