MCPcopy Create free account
hub / github.com/microsoft/BitNet / gen_transform_code

Function gen_transform_code

utils/codegen_tl2.py:626–674  ·  view source on GitHub ↗
(kernel_shapes)

Source from the content-addressed store, hash-verified

624 return kernel_code
625
626def gen_transform_code(kernel_shapes):
627 kernel_code = "\n\
628void 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
676def get_three_k_two_k(K, bk):
677 bk_num = K // bk

Callers 1

codegen_tl2.pyFile · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected