| 3 | from configparser import ConfigParser |
| 4 | |
| 5 | def gen_ctor_code(): |
| 6 | kernel_code = "\n\ |
| 7 | #include \"ggml-bitnet.h\"\n\ |
| 8 | #include <cstring>\n\ |
| 9 | #include <immintrin.h>\n\ |
| 10 | #define GGML_BITNET_MAX_NODES 8192\n\ |
| 11 | static bool initialized = false;\n\ |
| 12 | static bitnet_tensor_extra * bitnet_tensor_extras = nullptr;\n\ |
| 13 | static size_t bitnet_tensor_extras_index = 0;\n\ |
| 14 | static void * aligned_malloc(size_t size) {\n\ |
| 15 | #if defined(_WIN32)\n\ |
| 16 | return _aligned_malloc(size, 64);\n\ |
| 17 | #else\n\ |
| 18 | void * ptr = nullptr;\n\ |
| 19 | posix_memalign(&ptr, 64, size);\n\ |
| 20 | return ptr;\n\ |
| 21 | #endif\n\ |
| 22 | }\n\ |
| 23 | \n\ |
| 24 | static void aligned_free(void * ptr) {\n\ |
| 25 | #if defined(_WIN32)\n\ |
| 26 | _aligned_free(ptr);\n\ |
| 27 | #else\n\ |
| 28 | free(ptr);\n\ |
| 29 | #endif\n\ |
| 30 | }\n\ |
| 31 | #define BK2 32\n\ |
| 32 | #if defined __AVX2__\n\ |
| 33 | inline void _mm256_merge_epi32(const __m256i v0, const __m256i v1, __m256i *vl, __m256i *vh)\n\ |
| 34 | {\n\ |
| 35 | __m256i va = _mm256_permute4x64_epi64(v0, _MM_SHUFFLE(3, 1, 2, 0));\n\ |
| 36 | __m256i vb = _mm256_permute4x64_epi64(v1, _MM_SHUFFLE(3, 1, 2, 0));\n\ |
| 37 | *vl = _mm256_unpacklo_epi32(va, vb);\n\ |
| 38 | *vh = _mm256_unpackhi_epi32(va, vb);\n\ |
| 39 | }\n\ |
| 40 | inline void _mm256_merge_epi64(const __m256i v0, const __m256i v1, __m256i *vl, __m256i *vh)\n\ |
| 41 | {\n\ |
| 42 | __m256i va = _mm256_permute4x64_epi64(v0, _MM_SHUFFLE(3, 1, 2, 0));\n\ |
| 43 | __m256i vb = _mm256_permute4x64_epi64(v1, _MM_SHUFFLE(3, 1, 2, 0));\n\ |
| 44 | *vl = _mm256_unpacklo_epi64(va, vb);\n\ |
| 45 | *vh = _mm256_unpackhi_epi64(va, vb);\n\ |
| 46 | }\n\ |
| 47 | inline void _mm256_merge_si128(const __m256i v0, const __m256i v1, __m256i *vl, __m256i *vh)\n\ |
| 48 | {\n\ |
| 49 | *vl = _mm256_permute2x128_si256(v0, v1, _MM_SHUFFLE(0, 2, 0, 0));\n\ |
| 50 | *vh = _mm256_permute2x128_si256(v0, v1, _MM_SHUFFLE(0, 3, 0, 1));\n\ |
| 51 | }\n\ |
| 52 | inline void Transpose_8_8(\n\ |
| 53 | __m256i *v0,\n\ |
| 54 | __m256i *v1,\n\ |
| 55 | __m256i *v2,\n\ |
| 56 | __m256i *v3,\n\ |
| 57 | __m256i *v4,\n\ |
| 58 | __m256i *v5,\n\ |
| 59 | __m256i *v6,\n\ |
| 60 | __m256i *v7)\n\ |
| 61 | {\n\ |
| 62 | __m256i w0, w1, w2, w3, w4, w5, w6, w7;\n\ |