22 {
23 if (!x || !vy || k <= 0) {
24 return;
25 }
26 assert(k %
QK_K == 0);
27
28 const int nb = k /
QK_K;
30
31 for (int i = 0; i < nb; ++i) {
32
33 float max = 0.0f;
34 float amax = 0.0f;
35 for (
int j = 0; j <
QK_K; ++j) {
36 const float xv = x[j];
37 const float ax = fabsf(xv);
38 if (ax > amax) {
39 amax = ax;
40 max = xv;
41 }
42 }
43
44 if (amax == 0.0f) {
46 memset(y[i].qs, 0, sizeof(y[i].qs));
47 memset(y[i].bsums, 0, sizeof(y[i].bsums));
49 continue;
50 }
51
52 const float iscale = -127.0f / max;
53 const __m128 v_iscale = _mm_set1_ps(iscale);
54 const __m128 v_magic = _mm_set1_ps(12582912.0f);
55 const __m128i v_mantissa = _mm_set1_epi32(0x007fffff);
56 const __m128i v_bias = _mm_set1_epi32(0x00400000);
57 const __m128i v_min = _mm_set1_epi32(-128);
58 const __m128i v_max = _mm_set1_epi32(127);
59
60 for (
int j = 0; j <
QK_K; j += 16) {
61 const __m128 x0 = _mm_loadu_ps(x + j + 0);
62 const __m128 x1 = _mm_loadu_ps(x + j + 4);
63 const __m128 x2 = _mm_loadu_ps(x + j + 8);
64 const __m128 x3 = _mm_loadu_ps(x + j + 12);
65
66 __m128i q0 = _mm_sub_epi32(
67 _mm_and_si128(_mm_castps_si128(_mm_add_ps(_mm_mul_ps(x0, v_iscale), v_magic)), v_mantissa),
68 v_bias);
69 __m128i q1 = _mm_sub_epi32(
70 _mm_and_si128(_mm_castps_si128(_mm_add_ps(_mm_mul_ps(x1, v_iscale), v_magic)), v_mantissa),
71 v_bias);
72 __m128i q2 = _mm_sub_epi32(
73 _mm_and_si128(_mm_castps_si128(_mm_add_ps(_mm_mul_ps(x2, v_iscale), v_magic)), v_mantissa),
74 v_bias);
75 __m128i q3 = _mm_sub_epi32(
76 _mm_and_si128(_mm_castps_si128(_mm_add_ps(_mm_mul_ps(x3, v_iscale), v_magic)), v_mantissa),
77 v_bias);
78
79 q0 = _mm_min_epi32(_mm_max_epi32(q0, v_min), v_max);
80 q1 = _mm_min_epi32(_mm_max_epi32(q1, v_min), v_max);
81 q2 = _mm_min_epi32(_mm_max_epi32(q2, v_min), v_max);
82 q3 = _mm_min_epi32(_mm_max_epi32(q3, v_min), v_max);
83
84 const __m128i q01 = _mm_packs_epi32(q0, q1);
85 const __m128i q23 = _mm_packs_epi32(q2, q3);
86 const __m128i q0123 = _mm_packs_epi16(q01, q23);
87
88 _mm_storeu_si128((__m128i *)(y[i].qs + j), q0123);
89
90 int sum = 0;
91 for (int ii = 0; ii < 16; ++ii) {
92 sum += y[i].
qs[j + ii];
93 }
94 y[i].
bsums[j / 16] = (int16_t)sum;
95 }
96
97 y[i].
d = 1.0f / iscale;
99 }
100}