Revert "Slightly better CPU performance for SWA models (#1496)" (#1503)

This reverts commit f4125e8b1f.
This commit is contained in:
Kawrakow 2026-03-25 07:20:34 +01:00 committed by GitHub
parent 233225db8f
commit 5451b149d4
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
1 changed files with 22 additions and 31 deletions

View File

@ -1364,14 +1364,12 @@ void compute_helper(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q, in
for (; ik >=0 && Mc[ik] != 0; ik -= k_step);
ik += k_step;
for (int k1 = 0; k1 < ik/k_step; ++k1) {
if (k1 == ik/k_step-1 || Mc[k_step-1] == 0) {
#ifdef __aarch64__
KQHelper::multiply_mask_kq(kh, Dk, stride_m, q_f16, mr, fms);
KQHelper::multiply_mask_kq(kh, Dk, stride_m, q_f16, mr, fms);
#else
KQHelper::multiply_mask_kq(kh, stride_q, stride_m, q, mr, fms);
KQHelper::multiply_mask_kq(kh, stride_q, stride_m, q, mr, fms);
#endif
fqkv.accumulate_qkv(vh, fms);
}
fqkv.accumulate_qkv(vh, fms);
kh.next_block(k_step);
vh.next_block(k_step);
mr += k_step*sizeof(ggml_half);
@ -1431,12 +1429,9 @@ void compute_helper_q(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q,
for (; ik >=0 && Mc[ik] != 0; ik -= k_step);
ik += k_step;
for (int k1 = 0; k1 < ik/k_step; ++k1) {
Mc = (const uint16_t *)mr;
if (k1 == ik/k_step-1 || Mc[k_step-1] == 0) {
HelperQ80R8<Dk>::repack(k_step, kh.block, kh.stride, q8r8);
KQHelper::mul_mask_kq(khr8, stride_m, q8r, mr, fms);
fqkv.accumulate_qkv(vh, fms);
}
HelperQ80R8<Dk>::repack(k_step, kh.block, kh.stride, q8r8);
KQHelper::mul_mask_kq(khr8, stride_m, q8r, mr, fms);
fqkv.accumulate_qkv(vh, fms);
kh.next_block(k_step);
vh.next_block(k_step);
mr += k_step*sizeof(ggml_half);
@ -1465,20 +1460,17 @@ void compute_helper_q(KHelper& kh, VHelper& vh, int nq1, int nk1, int stride_q,
for (; ik >=0 && Mc[ik] != 0; ik -= k_step);
ik += k_step;
for (int k1 = 0; k1 < ik/k_step; ++k1) {
Mc = (const uint16_t *)mr;
if (k1 == ik/k_step-1 || Mc[k_step-1] == 0) {
#if FA_TIMING
t1 = Perf::cur_time();
KQHelper::mul_mask_kq(kh, stride_m, q8, mr, fms);
perf.accum_nolock(1, t1);
t1 = Perf::cur_time();
fqkv.accumulate_qkv(vh, fms);
perf.accum_nolock(2, t1);
t1 = Perf::cur_time();
KQHelper::mul_mask_kq(kh, stride_m, q8, mr, fms);
perf.accum_nolock(1, t1);
t1 = Perf::cur_time();
fqkv.accumulate_qkv(vh, fms);
perf.accum_nolock(2, t1);
#else
KQHelper::mul_mask_kq(kh, stride_m, q8, mr, fms);
fqkv.accumulate_qkv(vh, fms);
KQHelper::mul_mask_kq(kh, stride_m, q8, mr, fms);
fqkv.accumulate_qkv(vh, fms);
#endif
}
kh.next_block(k_step);
vh.next_block(k_step);
mr += k_step*sizeof(ggml_half);
@ -1990,18 +1982,17 @@ struct FlashAttnBF16 {
for (; ik >=0 && Mc[ik] != 0; ik -= k_step);
ik += k_step;
for (int k1 = 0; k1 < ik/k_step; ++k1) {
Mc = (const uint16_t *)mr;
if (k1 == ik/k_step-1 || Mc[k_step-1] == 0) {
#if FA_TIMING
FlashQKbf16<Dk, q_step, k_step>::multiply_mask_kq(kh, stride_m, q_bf16, mr, fms, perf);
t1 = Perf::cur_time();
fqkv.accumulate_qkv(vh, fms);
perf.accum_nolock(3, t1);
//t1 = Perf::cur_time();
FlashQKbf16<Dk, q_step, k_step>::multiply_mask_kq(kh, stride_m, q_bf16, mr, fms, perf);
//perf.accum_nolock(1, t1);
t1 = Perf::cur_time();
fqkv.accumulate_qkv(vh, fms);
perf.accum_nolock(3, t1);
#else
FlashQKbf16<Dk, q_step, k_step>::multiply_mask_kq(kh, stride_m, q_bf16, mr, fms);
fqkv.accumulate_qkv(vh, fms);
FlashQKbf16<Dk, q_step, k_step>::multiply_mask_kq(kh, stride_m, q_bf16, mr, fms);
fqkv.accumulate_qkv(vh, fms);
#endif
}
kh.next_block(k_step);
vh.next_block(k_step);
mr += k_step*sizeof(ggml_half);