87 if (!y || !x || !gamma || !W_q4k || M <= 0 || K <= 0) {
91 assert(K %
QK_K == 0);
92 const int nb = K /
QK_K;
97 assert(nb <= 32 &&
"K too large for stack buffer");
104#if defined(__AVX512F__)
106 __m512 sum_sq_vec = _mm512_setzero_ps();
108 for (; d + 16 <= K; d += 16) {
109 __m512 xv = _mm512_loadu_ps(&x[d]);
110 sum_sq_vec = _mm512_fmadd_ps(xv, xv, sum_sq_vec);
112 float sum_sq = _mm512_reduce_add_ps(sum_sq_vec);
114 sum_sq += x[d] * x[d];
117#elif defined(__AVX__)
119 __m256 sum_sq_vec = _mm256_setzero_ps();
121 for (; d + 8 <= K; d += 8) {
122 __m256 xv = _mm256_loadu_ps(&x[d]);
123 __m256 xv_sq = _mm256_mul_ps(xv, xv);
124 sum_sq_vec = _mm256_add_ps(sum_sq_vec, xv_sq);
126 float sum_sq = hsum256_ps_fused(sum_sq_vec);
128 sum_sq += x[d] * x[d];
134 for (
int d = 0; d < K; ++d) {
135 double v = (double)x[d];
140 float mean_sq = (float)sum_sq / (
float)K;
141 float rstd = 1.0f / sqrtf(mean_sq + eps);
148 for (
int i = 0; i < nb; ++i) {
149 const float *x_block = x + i *
QK_K;
150 const float *g_block = gamma + i *
QK_K;
153 float max_val = 0.0f;
156#if defined(__AVX512F__)
157 __m512 rstd_vec = _mm512_set1_ps(rstd);
158 __m512 max_vec = _mm512_setzero_ps();
159 __m512 sign_mask = _mm512_set1_ps(-0.0f);
161 for (
int j = 0; j <
QK_K; j += 16) {
162 __m512 xv = _mm512_loadu_ps(&x_block[j]);
163 __m512 gv = _mm512_loadu_ps(&g_block[j]);
164 __m512 norm = _mm512_mul_ps(_mm512_mul_ps(xv, rstd_vec), gv);
165 __m512 abs_norm = _mm512_andnot_ps(sign_mask, norm);
166 max_vec = _mm512_max_ps(max_vec, abs_norm);
169 __mmask16 gt_mask = _mm512_cmp_ps_mask(abs_norm, _mm512_set1_ps(amax), _CMP_GT_OQ);
171 float temp_amax = _mm512_reduce_max_ps(abs_norm);
172 if (temp_amax > amax) {
175 for (
int k = 0; k < 16; ++k) {
176 float v = x_block[j + k] * rstd * g_block[j + k];
177 if (fabsf(v) >= amax - 1e-6f) {
185 amax = _mm512_reduce_max_ps(max_vec);
187#elif defined(__AVX__)
188 __m256 rstd_vec = _mm256_set1_ps(rstd);
190 for (
int j = 0; j <
QK_K; j += 8) {
191 __m256 xv = _mm256_loadu_ps(&x_block[j]);
192 __m256 gv = _mm256_loadu_ps(&g_block[j]);
193 __m256 norm = _mm256_mul_ps(_mm256_mul_ps(xv, rstd_vec), gv);
197 _mm256_storeu_ps(norm_arr, norm);
198 for (
int k = 0; k < 8; ++k) {
199 float av = fabsf(norm_arr[k]);
202 max_val = norm_arr[k];
208 for (
int j = 0; j <
QK_K; ++j) {
209 float norm = x_block[j] * rstd * g_block[j];
210 float av = fabsf(norm);
220 q8_buffer[i].
d = 0.0f;
221 memset(q8_buffer[i].qs, 0,
sizeof(q8_buffer[i].qs));
222 memset(q8_buffer[i].bsums, 0,
sizeof(q8_buffer[i].bsums));
227 const float iscale = -127.0f / max_val;
228 q8_buffer[i].
d = 1.0f / iscale;
231 for (
int j = 0; j <
QK_K; ++j) {
232 float norm = x_block[j] * rstd * g_block[j];
234 v = (v > 127) ? 127 : ((v < -128) ? -128 : v);
235 q8_buffer[i].
qs[j] = (int8_t)v;
239 for (
int j = 0; j <
QK_K / 16; ++j) {
241 const int8_t *qs = &q8_buffer[i].
qs[j * 16];
242 for (
int k = 0; k < 16; ++k) {
245 q8_buffer[i].
bsums[j] = (int16_t)sum;
270 if (!y || !x || !gamma || !W_q4k || M <= 0 || K <= 0) {
274 assert(K %
QK_K == 0);
278 if (K > 4096)
return;
280 float norm_out[4096];
285 for (
int d = 0; d < K; ++d) {
286 sum_sq += (double)x[d] * (
double)x[d];
288 float rstd = 1.0f / sqrtf((
float)(sum_sq / K) + eps);
290 for (
int d = 0; d < K; ++d) {
291 norm_out[d] = x[d] * rstd * gamma[d];
void unfused_rmsnorm_linear_q4k_ref(float *y, const float *x, const float *gamma, const void *W_q4k, int M, int K, float eps)
Reference (unfused) implementation for correctness testing.
void fused_rmsnorm_linear_q4k(float *y, const float *x, const float *gamma, const void *W_q4k, int M, int K, float eps)
Fused RMSNorm + Q4_K Linear projection.