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

Function gen_tbl_impl

utils/codegen_tl2.py:279–530  ·  view source on GitHub ↗
(pre, BM, BK, bm, k_list)

Source from the content-addressed store, hash-verified

277 return kernel_code
278
279def 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\
286template<int batch_size, int K3>\n\
287inline 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\

Callers 1

codegen_tl2.pyFile · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected