41#if defined(__AVX__) || defined(__AVX2__) || defined(__AVXVNNI__)
52 int aligned_embed_dim,
61 int aligned_embed_dim,
70 int aligned_embed_dim,
79 int aligned_embed_dim,
88 int aligned_embed_dim,
94 const float *rstd_cache,
99 int aligned_embed_dim);
118 const char *forced = getenv(
"CK_QK_NORM_BACKWARD_ISA");
119 if (!forced || forced[0] ==
'\0' || strcmp(forced,
"auto") == 0) {
122 if (strcmp(forced,
"scalar") == 0) {
125 if (strcmp(forced,
"avx") == 0) {
128 if (strcmp(forced,
"avx2") == 0) {
131 if (strcmp(forced,
"avx_vnni") == 0) {
151#if defined(__AVXVNNI__)
167#if defined(__AVXVNNI__)
169#elif defined(__AVX2__)
171#elif defined(__AVX__)
184 for (
int r = 0; r < rows; ++r) {
185 const float *x = input + (size_t)r * (
size_t)head_dim;
187 for (
int d = 0; d < head_dim; ++d) {
188 double v = (double)x[d];
191 float mean_sq = (float)(sum_sq / (
double)head_dim);
192 rstd_cache[r] = 1.0f / sqrtf(mean_sq + eps);
197static inline float qk_norm_hsum256(__m256 v)
199 __m128 hi = _mm256_extractf128_ps(v, 1);
200 __m128 lo = _mm256_castps256_ps128(v);
201 __m128 sum = _mm_add_ps(lo, hi);
202 sum = _mm_hadd_ps(sum, sum);
203 sum = _mm_hadd_ps(sum, sum);
204 return _mm_cvtss_f32(sum);
207static void qk_norm_compute_rstd_avx(
const float *input,
213 for (
int r = 0; r < rows; ++r) {
214 const float *x = input + (size_t)r * (
size_t)head_dim;
215 __m256 sum_sq_v = _mm256_setzero_ps();
217 for (; d + 8 <= head_dim; d += 8) {
218 __m256 xv = _mm256_loadu_ps(&x[d]);
219 __m256 xv2 = _mm256_mul_ps(xv, xv);
220 sum_sq_v = _mm256_add_ps(sum_sq_v, xv2);
222 float sum_sq = qk_norm_hsum256(sum_sq_v);
223 for (; d < head_dim; ++d) {
224 sum_sq += x[d] * x[d];
226 float mean_sq = sum_sq / (float)head_dim;
227 rstd_cache[r] = 1.0f / sqrtf(mean_sq + eps);
233static void qk_norm_compute_rstd_avx2(
const float *input,
239 for (
int r = 0; r < rows; ++r) {
240 const float *x = input + (size_t)r * (
size_t)head_dim;
241 __m256 sum_sq_v = _mm256_setzero_ps();
243 for (; d + 8 <= head_dim; d += 8) {
244 __m256 xv = _mm256_loadu_ps(&x[d]);
246 sum_sq_v = _mm256_fmadd_ps(xv, xv, sum_sq_v);
248 __m256 xv2 = _mm256_mul_ps(xv, xv);
249 sum_sq_v = _mm256_add_ps(sum_sq_v, xv2);
252 float sum_sq = qk_norm_hsum256(sum_sq_v);
253 for (; d < head_dim; ++d) {
254 sum_sq += x[d] * x[d];
256 float mean_sq = sum_sq / (float)head_dim;
257 rstd_cache[r] = 1.0f / sqrtf(mean_sq + eps);
262#if defined(__AVXVNNI__)
263static void qk_norm_compute_rstd_avx_vnni(
const float *input,
270 qk_norm_compute_rstd_avx2(input, rstd_cache, rows, head_dim, eps);
271#elif defined(__AVX__)
272 qk_norm_compute_rstd_avx(input, rstd_cache, rows, head_dim, eps);
288#if defined(__AVXVNNI__)
290 qk_norm_compute_rstd_avx_vnni(input, rstd_cache, rows, head_dim, eps);
295 qk_norm_compute_rstd_avx2(input, rstd_cache, rows, head_dim, eps);
300 qk_norm_compute_rstd_avx(input, rstd_cache, rows, head_dim, eps);
327 const float *q_gamma,
const float *k_gamma,
328 int num_heads,
int num_kv_heads,
329 int num_tokens,
int head_dim,
float eps)
334 num_heads * num_tokens, head_dim, head_dim, eps);
339 num_kv_heads * num_tokens, head_dim, head_dim, eps);
343 const float *q_gamma,
const float *k_gamma,
344 int num_heads,
int num_kv_heads,
345 int num_tokens,
int head_dim,
float eps)
348 num_heads * num_tokens, head_dim, head_dim, eps);
350 num_kv_heads * num_tokens, head_dim, head_dim, eps);
354 const float *q_gamma,
const float *k_gamma,
355 int num_heads,
int num_kv_heads,
356 int num_tokens,
int head_dim,
float eps)
360 num_heads * num_tokens, head_dim, head_dim, eps);
363 num_kv_heads * num_tokens, head_dim, head_dim, eps);
367 const float *q_gamma,
368 const float *k_gamma,
369 int num_heads,
int num_kv_heads,
370 int num_tokens,
int head_dim,
375 num_heads * num_tokens, head_dim, head_dim, eps);
378 num_kv_heads * num_tokens, head_dim, head_dim, eps);
382 const float *q_gamma,
383 const float *k_gamma,
392 num_heads * num_tokens, head_dim, head_dim, eps);
395 num_kv_heads * num_tokens, head_dim, head_dim, eps);
406 const float *q_gamma,
413 num_heads * num_tokens, head_dim, head_dim, eps);
428 const float *q_in,
const float *k_in,
429 const float *q_gamma,
const float *k_gamma,
430 float *d_q_in,
float *d_k_in,
431 float *d_q_gamma,
float *d_k_gamma,
432 int num_heads,
int num_kv_heads,
433 int num_tokens,
int head_dim,
float eps)
435 int q_rows = num_heads * num_tokens;
436 int k_rows = num_kv_heads * num_tokens;
439 float q_rstd_cache[q_rows];
442 d_q_in, d_q_gamma, q_rows, head_dim, head_dim);
446 float k_rstd_cache[k_rows];
449 d_k_in, d_k_gamma, k_rows, head_dim, head_dim);
static int qk_norm_isa_compiled(QKNormISA isa)
static void qk_norm_compute_rstd_scalar(const float *input, float *rstd_cache, int rows, int head_dim, float eps)
void rmsnorm_forward_fp64_sum(const float *input, const float *gamma, float *output, float *rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)
void qk_norm_forward_qwen4_pytorch_bf16_storage(float *q, float *k, const float *q_gamma, const float *k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)
void rmsnorm_forward_qwen3next_pytorch_bf16_storage(const float *input, const float *gamma, float *output, float *rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)
void rmsnorm_forward_llama_production(const float *input, const float *gamma, float *output, float *rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)
static void qk_norm_compute_rstd(const float *input, float *rstd_cache, int rows, int head_dim, float eps)
void qk_norm_backward(const float *d_q_out, const float *d_k_out, const float *q_in, const float *k_in, const float *q_gamma, const float *k_gamma, float *d_q_in, float *d_k_in, float *d_q_gamma, float *d_k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)
void qk_norm_forward_llama_production(float *q, float *k, const float *q_gamma, const float *k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)
void qk_norm_forward_pytorch_bf16_storage(float *q, float *k, const float *q_gamma, const float *k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)
void qk_norm_forward(float *q, float *k, const float *q_gamma, const float *k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)
void qk_norm_forward_fp64_sum(float *q, float *k, const float *q_gamma, const float *k_gamma, int num_heads, int num_kv_heads, int num_tokens, int head_dim, float eps)
void rmsnorm_forward_pytorch_bf16_storage(const float *input, const float *gamma, float *output, float *rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)
void rmsnorm_forward(const float *input, const float *gamma, float *output, float *rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)
static int g_qk_norm_last_isa
void q_norm_forward(float *q, const float *q_gamma, int num_heads, int num_tokens, int head_dim, float eps)
int qk_norm_backward_last_isa(void)
void rmsnorm_backward(const float *d_output, const float *input, const float *gamma, const float *rstd_cache, float *d_input, float *d_gamma, int tokens, int d_model, int aligned_embed_dim)
static QKNormISA qk_norm_parse_forced_isa(void)
static QKNormISA qk_norm_select_isa(void)