| 277 | return kernel_code |
| 278 | |
| 279 | def gen_tbl_impl(pre, BM, BK, bm, k_list): |
| 280 | |
| 281 | kernel_code = "\ |
| 282 | #include <immintrin.h>\n\ |
| 283 | \n\ |
| 284 | #define BM{0} {1}\n\ |
| 285 | #define BBK{0} {2}\n\ |
| 286 | template<int batch_size, int K3>\n\ |
| 287 | inline void three_tbl_impl_{0}(int32_t* c, int8_t* lut, uint8_t* a, uint8_t* sign) {{\n\ |
| 288 | ".format(pre, BM, BK) |
| 289 | |
| 290 | kernel_code = "".join([kernel_code, "\ |
| 291 | #ifdef __AVX2__\n\ |
| 292 | const __m256i vec_mask = _mm256_set1_epi8(0x0f);\n\ |
| 293 | const __m256i vec_sign_mask = _mm256_set1_epi16(0x8000);\n\ |
| 294 | const __m256i vec_zero = _mm256_set1_epi8(0x00);\n\ |
| 295 | const __m256i vec_one = _mm256_set1_epi8(0xff);\n\ |
| 296 | const int KK = BBK{0} / 3;\n\ |
| 297 | #pragma unroll\n\ |
| 298 | for (int i = 0; i < BM{0}; i += 32) {{\n\ |
| 299 | __m256i vec_as[KK / 2];\n\ |
| 300 | __m256i vec_signs[KK / 8];\n\ |
| 301 | #pragma unroll\n\ |
| 302 | for (int ai = 0; ai < KK / 2; ai++) {{\n\ |
| 303 | vec_as[ai] = _mm256_loadu_si256(reinterpret_cast<__m256i*>(a + i * KK / 2 + ai * 32));\n\ |
| 304 | }}\n\ |
| 305 | #pragma unroll\n\ |
| 306 | for (int as = 0; as < KK / 8; as++) {{\n\ |
| 307 | vec_signs[as] = _mm256_loadu_si256(reinterpret_cast<__m256i*>(sign + i * KK / 8 + as * 32));\n\ |
| 308 | }}\n\ |
| 309 | #pragma unroll\n\ |
| 310 | for (int bs = 0; bs < batch_size; bs++) {{\n\ |
| 311 | __m256i vec_c0 = _mm256_setzero_si256();\n\ |
| 312 | __m256i vec_c1 = _mm256_setzero_si256();\n\ |
| 313 | #pragma unroll\n\ |
| 314 | for (int k = 0; k < KK / 8; k++) {{\n\ |
| 315 | __m256i vec_sign = vec_signs[k];\n\ |
| 316 | __m256i vec_a_0 = vec_as[k * 4 + 0];\n\ |
| 317 | __m128i vec_k1_0 = _mm_loadu_si128(reinterpret_cast<__m128i*>(lut + k * 32 * 8 + 0 * 64 + 0 + K3 / 3 * 32 * bs));\n\ |
| 318 | __m128i vec_k2_0 = _mm_loadu_si128(reinterpret_cast<__m128i*>(lut + k * 32 * 8 + 0 * 64 + 16 + K3 / 3 * 32 * bs));\n\ |
| 319 | __m128i vec_k3_0 = _mm_loadu_si128(reinterpret_cast<__m128i*>(lut + k * 32 * 8 + 0 * 64 + 32 + K3 / 3 * 32 * bs));\n\ |
| 320 | __m128i vec_k4_0 = _mm_loadu_si128(reinterpret_cast<__m128i*>(lut + k * 32 * 8 + 0 * 64 + 48 + K3 / 3 * 32 * bs));\n\ |
| 321 | __m256i vec_sign_left_hi_0 = _mm256_srai_epi16(_mm256_slli_epi16(vec_sign, (4 * 0)), 15);\n\ |
| 322 | __m256i vec_sign_left_lo_0 = _mm256_srai_epi16(_mm256_slli_epi16(vec_sign, (4 * 0 + 1)), 15);\n\ |
| 323 | __m256i vec_v_top_0 = _mm256_and_si256(_mm256_srli_epi16(vec_a_0, 4), vec_mask);\n\ |
| 324 | __m256i vec_v_top_fir_0 = _mm256_shuffle_epi8(_mm256_set_m128i(vec_k1_0, vec_k1_0), vec_v_top_0);\n\ |
| 325 | __m256i vec_v_top_sec_0 = _mm256_shuffle_epi8(_mm256_set_m128i(vec_k2_0, vec_k2_0), vec_v_top_0);\n\ |
| 326 | __m256i vec_sign_right_hi_0 = _mm256_srai_epi16(_mm256_slli_epi16(vec_sign, (4 * 0 + 2)), 15);\n\ |
| 327 | __m256i vec_sign_right_lo_0 = _mm256_srai_epi16(_mm256_slli_epi16(vec_sign, (4 * 0 + 3)), 15);\n\ |
| 328 | __m256i vec_v_bot_0 = _mm256_and_si256(vec_a_0, vec_mask);\n\ |
| 329 | __m256i vec_v_bot_fir_0 = _mm256_shuffle_epi8(_mm256_set_m128i(vec_k3_0, vec_k3_0), vec_v_bot_0);\n\ |
| 330 | __m256i vec_v_bot_sec_0 = _mm256_shuffle_epi8(_mm256_set_m128i(vec_k4_0, vec_k4_0), vec_v_bot_0);\n\ |
| 331 | __m256i vec_v_top_lo_0 = _mm256_xor_si256(_mm256_add_epi16(_mm256_unpackhi_epi8(vec_v_top_fir_0, vec_v_top_sec_0), vec_sign_left_lo_0), vec_sign_left_lo_0);\n\ |
| 332 | __m256i vec_v_top_hi_0 = _mm256_xor_si256(_mm256_add_epi16(_mm256_unpacklo_epi8(vec_v_top_fir_0, vec_v_top_sec_0), vec_sign_left_hi_0), vec_sign_left_hi_0);\n\ |
| 333 | __m256i vec_v_bot_lo_0 = _mm256_xor_si256(_mm256_add_epi16(_mm256_unpackhi_epi8(vec_v_bot_fir_0, vec_v_bot_sec_0), vec_sign_right_lo_0), vec_sign_right_lo_0);\n\ |
| 334 | __m256i vec_v_bot_hi_0 = _mm256_xor_si256(_mm256_add_epi16(_mm256_unpacklo_epi8(vec_v_bot_fir_0, vec_v_bot_sec_0), vec_sign_right_hi_0), vec_sign_right_hi_0);\n\ |
| 335 | vec_c0 = _mm256_add_epi16(vec_c0, vec_v_top_hi_0);\n\ |
| 336 | vec_c0 = _mm256_add_epi16(vec_c0, vec_v_bot_hi_0);\n\ |