91 const float *rms_weight,
117 for (
int j = 0; j < q_dim; j++) {
119 const float *wq_row = wq + j * embed_dim;
122 __m256 vsum = _mm256_setzero_ps();
124 for (; i + 7 < embed_dim; i += 8) {
125 __m256 vx = _mm256_loadu_ps(x + i);
126 __m256 vrms = _mm256_loadu_ps(rms_weight + i);
127 __m256 vw = _mm256_loadu_ps(wq_row + i);
128 __m256 vnormed = _mm256_mul_ps(vx, vrms);
129 vsum = _mm256_fmadd_ps(vw, vnormed, vsum);
132 __m128 vlow = _mm256_castps256_ps128(vsum);
133 __m128 vhigh = _mm256_extractf128_ps(vsum, 1);
134 vlow = _mm_add_ps(vlow, vhigh);
135 vlow = _mm_hadd_ps(vlow, vlow);
136 vlow = _mm_hadd_ps(vlow, vlow);
137 sum = _mm_cvtss_f32(vlow);
139 for (; i < embed_dim; i++) {
140 sum += wq_row[i] * x[i] * rms_weight[i];
143 for (
int i = 0; i < embed_dim; i++) {
144 sum += wq_row[i] * x[i] * rms_weight[i];
147 q_out[j] = sum * scale;
151 for (
int j = 0; j < kv_dim; j++) {
153 const float *wk_row = wk + j * embed_dim;
156 __m256 vsum = _mm256_setzero_ps();
158 for (; i + 7 < embed_dim; i += 8) {
159 __m256 vx = _mm256_loadu_ps(x + i);
160 __m256 vrms = _mm256_loadu_ps(rms_weight + i);
161 __m256 vw = _mm256_loadu_ps(wk_row + i);
162 __m256 vnormed = _mm256_mul_ps(vx, vrms);
163 vsum = _mm256_fmadd_ps(vw, vnormed, vsum);
165 __m128 vlow = _mm256_castps256_ps128(vsum);
166 __m128 vhigh = _mm256_extractf128_ps(vsum, 1);
167 vlow = _mm_add_ps(vlow, vhigh);
168 vlow = _mm_hadd_ps(vlow, vlow);
169 vlow = _mm_hadd_ps(vlow, vlow);
170 sum = _mm_cvtss_f32(vlow);
171 for (; i < embed_dim; i++) {
172 sum += wk_row[i] * x[i] * rms_weight[i];
175 for (
int i = 0; i < embed_dim; i++) {
176 sum += wk_row[i] * x[i] * rms_weight[i];
179 k_out[j] = sum * scale;
183 for (
int j = 0; j < kv_dim; j++) {
185 const float *wv_row = wv + j * embed_dim;
188 __m256 vsum = _mm256_setzero_ps();
190 for (; i + 7 < embed_dim; i += 8) {
191 __m256 vx = _mm256_loadu_ps(x + i);
192 __m256 vrms = _mm256_loadu_ps(rms_weight + i);
193 __m256 vw = _mm256_loadu_ps(wv_row + i);
194 __m256 vnormed = _mm256_mul_ps(vx, vrms);
195 vsum = _mm256_fmadd_ps(vw, vnormed, vsum);
197 __m128 vlow = _mm256_castps256_ps128(vsum);
198 __m128 vhigh = _mm256_extractf128_ps(vsum, 1);
199 vlow = _mm_add_ps(vlow, vhigh);
200 vlow = _mm_hadd_ps(vlow, vlow);
201 vlow = _mm_hadd_ps(vlow, vlow);
202 sum = _mm_cvtss_f32(vlow);
203 for (; i < embed_dim; i++) {
204 sum += wv_row[i] * x[i] * rms_weight[i];
207 for (
int i = 0; i < embed_dim; i++) {
208 sum += wv_row[i] * x[i] * rms_weight[i];
211 v_out[j] = sum * scale;
318 const float *rms_weight,
334 __m256 vscale = _mm256_set1_ps(scale);
347 for (
int j = 0; j < q_dim; j += 8) {
349 __m256 acc0 = _mm256_setzero_ps();
350 __m256 acc1 = _mm256_setzero_ps();
351 __m256 acc2 = _mm256_setzero_ps();
352 __m256 acc3 = _mm256_setzero_ps();
353 __m256 acc4 = _mm256_setzero_ps();
354 __m256 acc5 = _mm256_setzero_ps();
355 __m256 acc6 = _mm256_setzero_ps();
356 __m256 acc7 = _mm256_setzero_ps();
360 for (; i + 7 < embed_dim; i += 8) {
362 __m256 vx = _mm256_loadu_ps(x + i);
363 __m256 vrms = _mm256_loadu_ps(rms_weight + i);
364 __m256 normed = _mm256_mul_ps(_mm256_mul_ps(vx, vrms), vscale);
370 __m256 w0 = _mm256_loadu_ps(wq + (j+0)*embed_dim + i);
371 acc0 = _mm256_fmadd_ps(w0, normed, acc0);
374 __m256 w1 = _mm256_loadu_ps(wq + (j+1)*embed_dim + i);
375 acc1 = _mm256_fmadd_ps(w1, normed, acc1);
378 __m256 w2 = _mm256_loadu_ps(wq + (j+2)*embed_dim + i);
379 acc2 = _mm256_fmadd_ps(w2, normed, acc2);
382 __m256 w3 = _mm256_loadu_ps(wq + (j+3)*embed_dim + i);
383 acc3 = _mm256_fmadd_ps(w3, normed, acc3);
386 __m256 w4 = _mm256_loadu_ps(wq + (j+4)*embed_dim + i);
387 acc4 = _mm256_fmadd_ps(w4, normed, acc4);
390 __m256 w5 = _mm256_loadu_ps(wq + (j+5)*embed_dim + i);
391 acc5 = _mm256_fmadd_ps(w5, normed, acc5);
394 __m256 w6 = _mm256_loadu_ps(wq + (j+6)*embed_dim + i);
395 acc6 = _mm256_fmadd_ps(w6, normed, acc6);
398 __m256 w7 = _mm256_loadu_ps(wq + (j+7)*embed_dim + i);
399 acc7 = _mm256_fmadd_ps(w7, normed, acc7);
404 for (; i < embed_dim; i++) {
405 float normed_scalar = x[i] * rms_weight[i] * scale;
406 if (j + 0 < q_dim) acc0 = _mm256_add_ps(acc0, _mm256_set1_ps(wq[(j+0)*embed_dim + i] * normed_scalar));
407 if (j + 1 < q_dim) acc1 = _mm256_add_ps(acc1, _mm256_set1_ps(wq[(j+1)*embed_dim + i] * normed_scalar));
408 if (j + 2 < q_dim) acc2 = _mm256_add_ps(acc2, _mm256_set1_ps(wq[(j+2)*embed_dim + i] * normed_scalar));
409 if (j + 3 < q_dim) acc3 = _mm256_add_ps(acc3, _mm256_set1_ps(wq[(j+3)*embed_dim + i] * normed_scalar));
410 if (j + 4 < q_dim) acc4 = _mm256_add_ps(acc4, _mm256_set1_ps(wq[(j+4)*embed_dim + i] * normed_scalar));
411 if (j + 5 < q_dim) acc5 = _mm256_add_ps(acc5, _mm256_set1_ps(wq[(j+5)*embed_dim + i] * normed_scalar));
412 if (j + 6 < q_dim) acc6 = _mm256_add_ps(acc6, _mm256_set1_ps(wq[(j+6)*embed_dim + i] * normed_scalar));
413 if (j + 7 < q_dim) acc7 = _mm256_add_ps(acc7, _mm256_set1_ps(wq[(j+7)*embed_dim + i] * normed_scalar));
417 if (j + 0 < q_dim) q_out[j+0] = hsum256_ps(acc0);
418 if (j + 1 < q_dim) q_out[j+1] = hsum256_ps(acc1);
419 if (j + 2 < q_dim) q_out[j+2] = hsum256_ps(acc2);
420 if (j + 3 < q_dim) q_out[j+3] = hsum256_ps(acc3);
421 if (j + 4 < q_dim) q_out[j+4] = hsum256_ps(acc4);
422 if (j + 5 < q_dim) q_out[j+5] = hsum256_ps(acc5);
423 if (j + 6 < q_dim) q_out[j+6] = hsum256_ps(acc6);
424 if (j + 7 < q_dim) q_out[j+7] = hsum256_ps(acc7);
431 for (
int j = 0; j < kv_dim; j += 8) {
432 __m256 acc0 = _mm256_setzero_ps();
433 __m256 acc1 = _mm256_setzero_ps();
434 __m256 acc2 = _mm256_setzero_ps();
435 __m256 acc3 = _mm256_setzero_ps();
436 __m256 acc4 = _mm256_setzero_ps();
437 __m256 acc5 = _mm256_setzero_ps();
438 __m256 acc6 = _mm256_setzero_ps();
439 __m256 acc7 = _mm256_setzero_ps();
442 for (; i + 7 < embed_dim; i += 8) {
443 __m256 vx = _mm256_loadu_ps(x + i);
444 __m256 vrms = _mm256_loadu_ps(rms_weight + i);
445 __m256 normed = _mm256_mul_ps(_mm256_mul_ps(vx, vrms), vscale);
447 if (j + 0 < kv_dim) acc0 = _mm256_fmadd_ps(_mm256_loadu_ps(wk + (j+0)*embed_dim + i), normed, acc0);
448 if (j + 1 < kv_dim) acc1 = _mm256_fmadd_ps(_mm256_loadu_ps(wk + (j+1)*embed_dim + i), normed, acc1);
449 if (j + 2 < kv_dim) acc2 = _mm256_fmadd_ps(_mm256_loadu_ps(wk + (j+2)*embed_dim + i), normed, acc2);
450 if (j + 3 < kv_dim) acc3 = _mm256_fmadd_ps(_mm256_loadu_ps(wk + (j+3)*embed_dim + i), normed, acc3);
451 if (j + 4 < kv_dim) acc4 = _mm256_fmadd_ps(_mm256_loadu_ps(wk + (j+4)*embed_dim + i), normed, acc4);
452 if (j + 5 < kv_dim) acc5 = _mm256_fmadd_ps(_mm256_loadu_ps(wk + (j+5)*embed_dim + i), normed, acc5);
453 if (j + 6 < kv_dim) acc6 = _mm256_fmadd_ps(_mm256_loadu_ps(wk + (j+6)*embed_dim + i), normed, acc6);
454 if (j + 7 < kv_dim) acc7 = _mm256_fmadd_ps(_mm256_loadu_ps(wk + (j+7)*embed_dim + i), normed, acc7);
457 for (; i < embed_dim; i++) {
458 float normed_scalar = x[i] * rms_weight[i] * scale;
459 if (j + 0 < kv_dim) acc0 = _mm256_add_ps(acc0, _mm256_set1_ps(wk[(j+0)*embed_dim + i] * normed_scalar));
460 if (j + 1 < kv_dim) acc1 = _mm256_add_ps(acc1, _mm256_set1_ps(wk[(j+1)*embed_dim + i] * normed_scalar));
461 if (j + 2 < kv_dim) acc2 = _mm256_add_ps(acc2, _mm256_set1_ps(wk[(j+2)*embed_dim + i] * normed_scalar));
462 if (j + 3 < kv_dim) acc3 = _mm256_add_ps(acc3, _mm256_set1_ps(wk[(j+3)*embed_dim + i] * normed_scalar));
463 if (j + 4 < kv_dim) acc4 = _mm256_add_ps(acc4, _mm256_set1_ps(wk[(j+4)*embed_dim + i] * normed_scalar));
464 if (j + 5 < kv_dim) acc5 = _mm256_add_ps(acc5, _mm256_set1_ps(wk[(j+5)*embed_dim + i] * normed_scalar));
465 if (j + 6 < kv_dim) acc6 = _mm256_add_ps(acc6, _mm256_set1_ps(wk[(j+6)*embed_dim + i] * normed_scalar));
466 if (j + 7 < kv_dim) acc7 = _mm256_add_ps(acc7, _mm256_set1_ps(wk[(j+7)*embed_dim + i] * normed_scalar));
469 if (j + 0 < kv_dim) k_out[j+0] = hsum256_ps(acc0);
470 if (j + 1 < kv_dim) k_out[j+1] = hsum256_ps(acc1);
471 if (j + 2 < kv_dim) k_out[j+2] = hsum256_ps(acc2);
472 if (j + 3 < kv_dim) k_out[j+3] = hsum256_ps(acc3);
473 if (j + 4 < kv_dim) k_out[j+4] = hsum256_ps(acc4);
474 if (j + 5 < kv_dim) k_out[j+5] = hsum256_ps(acc5);
475 if (j + 6 < kv_dim) k_out[j+6] = hsum256_ps(acc6);
476 if (j + 7 < kv_dim) k_out[j+7] = hsum256_ps(acc7);
483 for (
int j = 0; j < kv_dim; j += 8) {
484 __m256 acc0 = _mm256_setzero_ps();
485 __m256 acc1 = _mm256_setzero_ps();
486 __m256 acc2 = _mm256_setzero_ps();
487 __m256 acc3 = _mm256_setzero_ps();
488 __m256 acc4 = _mm256_setzero_ps();
489 __m256 acc5 = _mm256_setzero_ps();
490 __m256 acc6 = _mm256_setzero_ps();
491 __m256 acc7 = _mm256_setzero_ps();
494 for (; i + 7 < embed_dim; i += 8) {
495 __m256 vx = _mm256_loadu_ps(x + i);
496 __m256 vrms = _mm256_loadu_ps(rms_weight + i);
497 __m256 normed = _mm256_mul_ps(_mm256_mul_ps(vx, vrms), vscale);
499 if (j + 0 < kv_dim) acc0 = _mm256_fmadd_ps(_mm256_loadu_ps(wv + (j+0)*embed_dim + i), normed, acc0);
500 if (j + 1 < kv_dim) acc1 = _mm256_fmadd_ps(_mm256_loadu_ps(wv + (j+1)*embed_dim + i), normed, acc1);
501 if (j + 2 < kv_dim) acc2 = _mm256_fmadd_ps(_mm256_loadu_ps(wv + (j+2)*embed_dim + i), normed, acc2);
502 if (j + 3 < kv_dim) acc3 = _mm256_fmadd_ps(_mm256_loadu_ps(wv + (j+3)*embed_dim + i), normed, acc3);
503 if (j + 4 < kv_dim) acc4 = _mm256_fmadd_ps(_mm256_loadu_ps(wv + (j+4)*embed_dim + i), normed, acc4);
504 if (j + 5 < kv_dim) acc5 = _mm256_fmadd_ps(_mm256_loadu_ps(wv + (j+5)*embed_dim + i), normed, acc5);
505 if (j + 6 < kv_dim) acc6 = _mm256_fmadd_ps(_mm256_loadu_ps(wv + (j+6)*embed_dim + i), normed, acc6);
506 if (j + 7 < kv_dim) acc7 = _mm256_fmadd_ps(_mm256_loadu_ps(wv + (j+7)*embed_dim + i), normed, acc7);
509 for (; i < embed_dim; i++) {
510 float normed_scalar = x[i] * rms_weight[i] * scale;
511 if (j + 0 < kv_dim) acc0 = _mm256_add_ps(acc0, _mm256_set1_ps(wv[(j+0)*embed_dim + i] * normed_scalar));
512 if (j + 1 < kv_dim) acc1 = _mm256_add_ps(acc1, _mm256_set1_ps(wv[(j+1)*embed_dim + i] * normed_scalar));
513 if (j + 2 < kv_dim) acc2 = _mm256_add_ps(acc2, _mm256_set1_ps(wv[(j+2)*embed_dim + i] * normed_scalar));
514 if (j + 3 < kv_dim) acc3 = _mm256_add_ps(acc3, _mm256_set1_ps(wv[(j+3)*embed_dim + i] * normed_scalar));
515 if (j + 4 < kv_dim) acc4 = _mm256_add_ps(acc4, _mm256_set1_ps(wv[(j+4)*embed_dim + i] * normed_scalar));
516 if (j + 5 < kv_dim) acc5 = _mm256_add_ps(acc5, _mm256_set1_ps(wv[(j+5)*embed_dim + i] * normed_scalar));
517 if (j + 6 < kv_dim) acc6 = _mm256_add_ps(acc6, _mm256_set1_ps(wv[(j+6)*embed_dim + i] * normed_scalar));
518 if (j + 7 < kv_dim) acc7 = _mm256_add_ps(acc7, _mm256_set1_ps(wv[(j+7)*embed_dim + i] * normed_scalar));
521 if (j + 0 < kv_dim) v_out[j+0] = hsum256_ps(acc0);
522 if (j + 1 < kv_dim) v_out[j+1] = hsum256_ps(acc1);
523 if (j + 2 < kv_dim) v_out[j+2] = hsum256_ps(acc2);
524 if (j + 3 < kv_dim) v_out[j+3] = hsum256_ps(acc3);
525 if (j + 4 < kv_dim) v_out[j+4] = hsum256_ps(acc4);
526 if (j + 5 < kv_dim) v_out[j+5] = hsum256_ps(acc5);
527 if (j + 6 < kv_dim) v_out[j+6] = hsum256_ps(acc6);
528 if (j + 7 < kv_dim) v_out[j+7] = hsum256_ps(acc7);
533 for (
int j = 0; j < q_dim; j++) {
535 for (
int i = 0; i < embed_dim; i++) {
536 float normed = x[i] * rms_weight[i] * scale;
537 sum += wq[j * embed_dim + i] * normed;
541 for (
int j = 0; j < kv_dim; j++) {
543 for (
int i = 0; i < embed_dim; i++) {
544 float normed = x[i] * rms_weight[i] * scale;
545 sum += wk[j * embed_dim + i] * normed;
549 for (
int j = 0; j < kv_dim; j++) {
551 for (
int i = 0; i < embed_dim; i++) {
552 float normed = x[i] * rms_weight[i] * scale;
553 sum += wv[j * embed_dim + i] * normed;
588 const float *rms_weight,
604 __m256 vscale = _mm256_set1_ps(scale);
610 for (
int j = 0; j < kv_dim; j++) {
611 __m256 q_acc = _mm256_setzero_ps();
612 __m256 k_acc = _mm256_setzero_ps();
613 __m256 v_acc = _mm256_setzero_ps();
616 for (; i + 7 < embed_dim; i += 8) {
618 __m256 vx = _mm256_loadu_ps(x + i);
619 __m256 vrms = _mm256_loadu_ps(rms_weight + i);
620 __m256 normed = _mm256_mul_ps(_mm256_mul_ps(vx, vrms), vscale);
623 __m256 wq_row = _mm256_loadu_ps(wq + j * embed_dim + i);
624 __m256 wk_row = _mm256_loadu_ps(wk + j * embed_dim + i);
625 __m256 wv_row = _mm256_loadu_ps(wv + j * embed_dim + i);
628 q_acc = _mm256_fmadd_ps(wq_row, normed, q_acc);
629 k_acc = _mm256_fmadd_ps(wk_row, normed, k_acc);
630 v_acc = _mm256_fmadd_ps(wv_row, normed, v_acc);
634 float q_sum = hsum256_ps(q_acc);
635 float k_sum = hsum256_ps(k_acc);
636 float v_sum = hsum256_ps(v_acc);
638 for (; i < embed_dim; i++) {
639 float normed = x[i] * rms_weight[i] * scale;
640 q_sum += wq[j * embed_dim + i] * normed;
641 k_sum += wk[j * embed_dim + i] * normed;
642 v_sum += wv[j * embed_dim + i] * normed;
654 for (
int j = kv_dim; j < q_dim; j++) {
655 __m256 q_acc = _mm256_setzero_ps();
658 for (; i + 7 < embed_dim; i += 8) {
659 __m256 vx = _mm256_loadu_ps(x + i);
660 __m256 vrms = _mm256_loadu_ps(rms_weight + i);
661 __m256 normed = _mm256_mul_ps(_mm256_mul_ps(vx, vrms), vscale);
663 __m256 wq_row = _mm256_loadu_ps(wq + j * embed_dim + i);
664 q_acc = _mm256_fmadd_ps(wq_row, normed, q_acc);
667 float q_sum = hsum256_ps(q_acc);
668 for (; i < embed_dim; i++) {
669 float normed = x[i] * rms_weight[i] * scale;
670 q_sum += wq[j * embed_dim + i] * normed;
678 for (
int j = 0; j < kv_dim; j++) {
679 float q_sum = 0.0f, k_sum = 0.0f, v_sum = 0.0f;
680 for (
int i = 0; i < embed_dim; i++) {
681 float normed = x[i] * rms_weight[i] * scale;
682 q_sum += wq[j * embed_dim + i] * normed;
683 k_sum += wk[j * embed_dim + i] * normed;
684 v_sum += wv[j * embed_dim + i] * normed;
690 for (
int j = kv_dim; j < q_dim; j++) {
692 for (
int i = 0; i < embed_dim; i++) {
693 float normed = x[i] * rms_weight[i] * scale;
694 q_sum += wq[j * embed_dim + i] * normed;