368 if (!grad || !weight || !m || !v || numel == 0) {
373 float bias_correction1 = 1.0f - powf(beta1, (
float)step);
374 float bias_correction2 = 1.0f - powf(beta2, (
float)step);
377 float one_minus_beta1 = 1.0f - beta1;
378 float one_minus_beta2 = 1.0f - beta2;
381 const double beta1_d = (double)beta1;
382 const double beta2_d = (double)beta2;
383 const double one_minus_beta1_d = 1.0 - beta1_d;
384 const double one_minus_beta2_d = 1.0 - beta2_d;
385 const double lr_d = (double)lr;
386 const double eps_d = (double)eps;
387 const double wd_d = (double)weight_decay;
389 const double bc1 = 1.0 - pow(beta1_d, (
double)step);
390 const double bc2 = 1.0 - pow(beta2_d, (
double)step);
391 const double step_size = lr_d / bc1;
392 const double bc2_sqrt = sqrt(bc2);
393 const double wd_scale = 1.0 - lr_d * wd_d;
395 for (
size_t i = 0; i < numel; ++i) {
396 double g = (double)grad[i];
397 double w = (double)weight[i];
398 double m_i = (double)m[i];
399 double v_i = (double)v[i];
401 m_i = beta1_d * m_i + one_minus_beta1_d * g;
402 v_i = beta2_d * v_i + one_minus_beta2_d * g * g;
407 double denom = sqrt(v_i) / bc2_sqrt + eps_d;
408 w -= step_size * (m_i / denom);
412 weight[i] = (float)w;
417#if defined(__AVX512F__)
419 __m512 v_beta1 = _mm512_set1_ps(beta1);
420 __m512 v_beta2 = _mm512_set1_ps(beta2);
421 __m512 v_one_minus_beta1 = _mm512_set1_ps(one_minus_beta1);
422 __m512 v_one_minus_beta2 = _mm512_set1_ps(one_minus_beta2);
423 __m512 v_lr = _mm512_set1_ps(lr);
424 __m512 v_eps = _mm512_set1_ps(eps);
425 __m512 v_weight_decay = _mm512_set1_ps(weight_decay);
426 __m512 v_bc1_inv = _mm512_set1_ps(1.0f / bias_correction1);
427 __m512 v_bc2_inv = _mm512_set1_ps(1.0f / bias_correction2);
430 for (; i + 16 <= numel; i += 16) {
431 __m512 g = _mm512_loadu_ps(&grad[i]);
432 __m512 w = _mm512_loadu_ps(&weight[i]);
433 __m512 m_val = _mm512_loadu_ps(&m[i]);
434 __m512 v_val = _mm512_loadu_ps(&v[i]);
437 m_val = _mm512_fmadd_ps(v_beta1, m_val, _mm512_mul_ps(v_one_minus_beta1, g));
440 __m512 g_sq = _mm512_mul_ps(g, g);
441 v_val = _mm512_fmadd_ps(v_beta2, v_val, _mm512_mul_ps(v_one_minus_beta2, g_sq));
444 __m512 m_hat = _mm512_mul_ps(m_val, v_bc1_inv);
445 __m512 v_hat = _mm512_mul_ps(v_val, v_bc2_inv);
448 __m512 denom = _mm512_add_ps(_mm512_sqrt_ps(v_hat), v_eps);
449 __m512 update = _mm512_div_ps(m_hat, denom);
450 update = _mm512_fmadd_ps(v_weight_decay, w, update);
451 w = _mm512_fnmadd_ps(v_lr, update, w);
453 _mm512_storeu_ps(&weight[i], w);
454 _mm512_storeu_ps(&m[i], m_val);
455 _mm512_storeu_ps(&v[i], v_val);
459 for (; i < numel; ++i) {
462 m[i] = beta1 * m[i] + one_minus_beta1 * g;
463 v[i] = beta2 * v[i] + one_minus_beta2 * g * g;
464 float m_hat = m[i] / bias_correction1;
465 float v_hat = v[i] / bias_correction2;
466 weight[i] = w - lr * (m_hat / (sqrtf(v_hat) + eps) + weight_decay * w);
469#elif defined(__AVX__)
471 __m256 v_beta1 = _mm256_set1_ps(beta1);
472 __m256 v_beta2 = _mm256_set1_ps(beta2);
473 __m256 v_one_minus_beta1 = _mm256_set1_ps(one_minus_beta1);
474 __m256 v_one_minus_beta2 = _mm256_set1_ps(one_minus_beta2);
475 __m256 v_lr = _mm256_set1_ps(lr);
476 __m256 v_eps = _mm256_set1_ps(eps);
477 __m256 v_weight_decay = _mm256_set1_ps(weight_decay);
478 __m256 v_bc1_inv = _mm256_set1_ps(1.0f / bias_correction1);
479 __m256 v_bc2_inv = _mm256_set1_ps(1.0f / bias_correction2);
482 for (; i + 8 <= numel; i += 8) {
483 __m256 g = _mm256_loadu_ps(&grad[i]);
484 __m256 w = _mm256_loadu_ps(&weight[i]);
485 __m256 m_val = _mm256_loadu_ps(&m[i]);
486 __m256 v_val = _mm256_loadu_ps(&v[i]);
489 m_val = _mm256_add_ps(_mm256_mul_ps(v_beta1, m_val),
490 _mm256_mul_ps(v_one_minus_beta1, g));
493 __m256 g_sq = _mm256_mul_ps(g, g);
494 v_val = _mm256_add_ps(_mm256_mul_ps(v_beta2, v_val),
495 _mm256_mul_ps(v_one_minus_beta2, g_sq));
498 __m256 m_hat = _mm256_mul_ps(m_val, v_bc1_inv);
499 __m256 v_hat = _mm256_mul_ps(v_val, v_bc2_inv);
502 __m256 denom = _mm256_add_ps(_mm256_sqrt_ps(v_hat), v_eps);
503 __m256 update = _mm256_div_ps(m_hat, denom);
504 update = _mm256_add_ps(update, _mm256_mul_ps(v_weight_decay, w));
505 w = _mm256_sub_ps(w, _mm256_mul_ps(v_lr, update));
507 _mm256_storeu_ps(&weight[i], w);
508 _mm256_storeu_ps(&m[i], m_val);
509 _mm256_storeu_ps(&v[i], v_val);
513 for (; i < numel; ++i) {
516 m[i] = beta1 * m[i] + one_minus_beta1 * g;
517 v[i] = beta2 * v[i] + one_minus_beta2 * g * g;
518 float m_hat = m[i] / bias_correction1;
519 float v_hat = v[i] / bias_correction2;
520 weight[i] = w - lr * (m_hat / (sqrtf(v_hat) + eps) + weight_decay * w);
523#elif defined(__SSE2__)
525 __m128 v_beta1 = _mm_set1_ps(beta1);
526 __m128 v_beta2 = _mm_set1_ps(beta2);
527 __m128 v_one_minus_beta1 = _mm_set1_ps(one_minus_beta1);
528 __m128 v_one_minus_beta2 = _mm_set1_ps(one_minus_beta2);
529 __m128 v_lr = _mm_set1_ps(lr);
530 __m128 v_eps = _mm_set1_ps(eps);
531 __m128 v_weight_decay = _mm_set1_ps(weight_decay);
532 __m128 v_bc1_inv = _mm_set1_ps(1.0f / bias_correction1);
533 __m128 v_bc2_inv = _mm_set1_ps(1.0f / bias_correction2);
536 for (; i + 4 <= numel; i += 4) {
537 __m128 g = _mm_loadu_ps(&grad[i]);
538 __m128 w = _mm_loadu_ps(&weight[i]);
539 __m128 m_val = _mm_loadu_ps(&m[i]);
540 __m128 v_val = _mm_loadu_ps(&v[i]);
543 m_val = _mm_add_ps(_mm_mul_ps(v_beta1, m_val),
544 _mm_mul_ps(v_one_minus_beta1, g));
547 __m128 g_sq = _mm_mul_ps(g, g);
548 v_val = _mm_add_ps(_mm_mul_ps(v_beta2, v_val),
549 _mm_mul_ps(v_one_minus_beta2, g_sq));
552 __m128 m_hat = _mm_mul_ps(m_val, v_bc1_inv);
553 __m128 v_hat = _mm_mul_ps(v_val, v_bc2_inv);
556 __m128 denom = _mm_add_ps(_mm_sqrt_ps(v_hat), v_eps);
557 __m128 update = _mm_div_ps(m_hat, denom);
558 update = _mm_add_ps(update, _mm_mul_ps(v_weight_decay, w));
559 w = _mm_sub_ps(w, _mm_mul_ps(v_lr, update));
561 _mm_storeu_ps(&weight[i], w);
562 _mm_storeu_ps(&m[i], m_val);
563 _mm_storeu_ps(&v[i], v_val);
567 for (; i < numel; ++i) {
570 m[i] = beta1 * m[i] + one_minus_beta1 * g;
571 v[i] = beta2 * v[i] + one_minus_beta2 * g * g;
572 float m_hat = m[i] / bias_correction1;
573 float v_hat = v[i] / bias_correction2;
574 weight[i] = w - lr * (m_hat / (sqrtf(v_hat) + eps) + weight_decay * w);
579 for (
size_t i = 0; i < numel; ++i) {
582 m[i] = beta1 * m[i] + one_minus_beta1 * g;
583 v[i] = beta2 * v[i] + one_minus_beta2 * g * g;
584 float m_hat = m[i] / bias_correction1;
585 float v_hat = v[i] / bias_correction2;
586 weight[i] = w - lr * (m_hat / (sqrtf(v_hat) + eps) + weight_decay * w);
640 float *
const *weights,
641 float *
const *m_states,
642 float *
const *v_states,
643 const size_t *numels,
653 if (!grads || !weights || !m_states || !v_states || !numels || tensor_count <= 0) {
657 size_t total_numel = 0;
658 int valid_tensors = 0;
659 for (
int i = 0; i < tensor_count; ++i) {
660 if (grads[i] && weights[i] && m_states[i] && v_states[i] && numels[i] > 0) {
661 total_numel += numels[i];
665 if (valid_tensors == 0 || total_numel == 0) {
669 float grad_scale = 1.0f;
670 if (max_grad_norm > 0.0f) {
672 if (global_norm > max_grad_norm) {
673 grad_scale = max_grad_norm / global_norm;
682 for (
int i = 0; i < tensor_count; ++i) {
684 float *w = weights[i];
685 float *m = m_states[i];
686 float *v = v_states[i];
687 size_t n = numels[i];
688 if (!g || !w || !m || !v || n == 0) {
691 if (grad_scale != 1.0f) {
694 adamw_update_f32_impl(g, w, m, v, n, lr, beta1, beta2, eps, weight_decay, step);
699 if (active_nth > valid_tensors) {
700 active_nth = valid_tensors;
702 if (active_nth <= 1) {
703 for (
int i = 0; i < tensor_count; ++i) {
705 float *w = weights[i];
706 float *m = m_states[i];
707 float *v = v_states[i];
708 size_t n = numels[i];
709 if (!g || !w || !m || !v || n == 0) {
712 if (grad_scale != 1.0f) {
715 adamw_update_f32_impl(g, w, m, v, n, lr, beta1, beta2, eps, weight_decay, step);
720 ck_adamw_multi_parallel_args_t args = {
723 .m_states = m_states,
724 .v_states = v_states,
726 .tensor_count = tensor_count,
731 .weight_decay = weight_decay,
732 .grad_scale = grad_scale,
762 if (!grad || !weight || !velocity || numel == 0) {
766#if defined(__AVX512F__)
768 __m512 v_lr = _mm512_set1_ps(lr);
769 __m512 v_momentum = _mm512_set1_ps(momentum);
770 __m512 v_weight_decay = _mm512_set1_ps(weight_decay);
773 for (; i + 16 <= numel; i += 16) {
774 __m512 g = _mm512_loadu_ps(&grad[i]);
775 __m512 w = _mm512_loadu_ps(&weight[i]);
776 __m512 vel = _mm512_loadu_ps(&velocity[i]);
778 vel = _mm512_fmadd_ps(v_momentum, vel, g);
779 __m512 update = _mm512_fmadd_ps(v_weight_decay, w, vel);
780 w = _mm512_fnmadd_ps(v_lr, update, w);
782 _mm512_storeu_ps(&weight[i], w);
783 _mm512_storeu_ps(&velocity[i], vel);
786 for (; i < numel; ++i) {
787 velocity[i] = momentum * velocity[i] + grad[i];
788 weight[i] = weight[i] - lr * (velocity[i] + weight_decay * weight[i]);
791#elif defined(__AVX__)
793 __m256 v_lr = _mm256_set1_ps(lr);
794 __m256 v_momentum = _mm256_set1_ps(momentum);
795 __m256 v_weight_decay = _mm256_set1_ps(weight_decay);
798 for (; i + 8 <= numel; i += 8) {
799 __m256 g = _mm256_loadu_ps(&grad[i]);
800 __m256 w = _mm256_loadu_ps(&weight[i]);
801 __m256 vel = _mm256_loadu_ps(&velocity[i]);
804 vel = _mm256_add_ps(_mm256_mul_ps(v_momentum, vel), g);
807 __m256 update = _mm256_add_ps(vel, _mm256_mul_ps(v_weight_decay, w));
808 w = _mm256_sub_ps(w, _mm256_mul_ps(v_lr, update));
810 _mm256_storeu_ps(&weight[i], w);
811 _mm256_storeu_ps(&velocity[i], vel);
814 for (; i < numel; ++i) {
815 velocity[i] = momentum * velocity[i] + grad[i];
816 weight[i] = weight[i] - lr * (velocity[i] + weight_decay * weight[i]);
819#elif defined(__SSE2__)
821 __m128 v_lr = _mm_set1_ps(lr);
822 __m128 v_momentum = _mm_set1_ps(momentum);
823 __m128 v_weight_decay = _mm_set1_ps(weight_decay);
826 for (; i + 4 <= numel; i += 4) {
827 __m128 g = _mm_loadu_ps(&grad[i]);
828 __m128 w = _mm_loadu_ps(&weight[i]);
829 __m128 vel = _mm_loadu_ps(&velocity[i]);
831 vel = _mm_add_ps(_mm_mul_ps(v_momentum, vel), g);
832 __m128 update = _mm_add_ps(vel, _mm_mul_ps(v_weight_decay, w));
833 w = _mm_sub_ps(w, _mm_mul_ps(v_lr, update));
835 _mm_storeu_ps(&weight[i], w);
836 _mm_storeu_ps(&velocity[i], vel);
839 for (; i < numel; ++i) {
840 velocity[i] = momentum * velocity[i] + grad[i];
841 weight[i] = weight[i] - lr * (velocity[i] + weight_decay * weight[i]);
846 for (
size_t i = 0; i < numel; ++i) {
847 velocity[i] = momentum * velocity[i] + grad[i];
848 weight[i] = weight[i] - lr * (velocity[i] + weight_decay * weight[i]);
1101 if (!grad || numel == 0) {
1105 double sum_sq = 0.0;
1106#if defined(__AVX512F__)
1107 __m512 acc = _mm512_setzero_ps();
1109 for (; i + 16 <= numel; i += 16) {
1110 __m512 g = _mm512_loadu_ps(&grad[i]);
1111 acc = _mm512_fmadd_ps(g, g, acc);
1113 sum_sq = _mm512_reduce_add_ps(acc);
1114 for (; i < numel; ++i) {
1115 sum_sq += (double)grad[i] * (
double)grad[i];
1117#elif defined(__AVX__)
1118 __m256 acc = _mm256_setzero_ps();
1120 for (; i + 8 <= numel; i += 8) {
1121 __m256 g = _mm256_loadu_ps(&grad[i]);
1122 acc = _mm256_add_ps(acc, _mm256_mul_ps(g, g));
1124 __m128 hi = _mm256_extractf128_ps(acc, 1);
1125 __m128 lo = _mm256_castps256_ps128(acc);
1126 __m128 sum4 = _mm_add_ps(lo, hi);
1127 __m128 shuf = _mm_movehdup_ps(sum4);
1128 __m128 sums = _mm_add_ps(sum4, shuf);
1129 shuf = _mm_movehl_ps(shuf, sums);
1130 sums = _mm_add_ss(sums, shuf);
1131 sum_sq = _mm_cvtss_f32(sums);
1132 for (; i < numel; ++i) {
1133 sum_sq += (double)grad[i] * (
double)grad[i];
1135#elif defined(__SSE2__)
1136 __m128 acc = _mm_setzero_ps();
1138 for (; i + 4 <= numel; i += 4) {
1139 __m128 g = _mm_loadu_ps(&grad[i]);
1140 acc = _mm_add_ps(acc, _mm_mul_ps(g, g));
1142 __m128 shuf = _mm_shuffle_ps(acc, acc, _MM_SHUFFLE(2, 3, 0, 1));
1143 __m128 sums = _mm_add_ps(acc, shuf);
1144 shuf = _mm_movehl_ps(shuf, sums);
1145 sums = _mm_add_ss(sums, shuf);
1146 sum_sq = _mm_cvtss_f32(sums);
1147 for (; i < numel; ++i) {
1148 sum_sq += (double)grad[i] * (
double)grad[i];
1151 for (
size_t i = 0; i < numel; ++i) {
1152 sum_sq += (double)grad[i] * (
double)grad[i];