Faster indexer top_k for very long context (CPU) (#2206)

* Faster indexer top_k for very long context (CPU)

* Minor
This commit is contained in:
Kawrakow
2026-08-01 09:17:51 +03:00
committed by GitHub
parent 8ba790e8ad
commit f2bde5749b
+62 -33
View File
@@ -1740,6 +1740,7 @@ bool iqk_fused_delta_net(int head_dim, int n_heads, int gqa_ratio, int repeat_ty
}
namespace {
constexpr int k_indexer_chunks = 64;
size_t iqk_idx_topk_work_wbs_per_thread(const struct ggml_tensor * dst, int nth) {
auto k = dst->src[0];
auto q = dst->src[1];
@@ -1750,7 +1751,7 @@ size_t iqk_idx_topk_work_wbs_per_thread(const struct ggml_tensor * dst, int nth)
auto row_size_q = ggml_row_size(tt.vec_dot_type, q->ne[0]);
size = row_size_q * q->ne[1];
}
size += k->ne[1] * q->ne[1] * sizeof(float);
size += k_indexer_chunks * q->ne[1] * sizeof(float);
size += k->ne[1] * sizeof(float);
size += k->ne[1] * sizeof(int32_t);
size = GGML_PAD(size, 128);
@@ -1831,7 +1832,7 @@ bool iqk_indexer_topk(struct ggml_tensor * dst, void * work_buffer, barrier_t ba
if (q->ne[2] >= nth) {
auto kq = (float *)(work + quantize_size);
auto score = kq + k->ne[1]*q->ne[1];
auto score = kq + k_indexer_chunks*q->ne[1];
auto sorted = (int32_t *)(score + k->ne[1]);
for (int iq = ith; iq < q->ne[2]; iq += nth) {
auto this_q = (const char *)q->data + iq*q->nb[2];
@@ -1841,22 +1842,48 @@ bool iqk_indexer_topk(struct ggml_tensor * dst, void * work_buffer, barrier_t ba
from_float((const float *)this_q, work, q->ne[0] * q->ne[1]);
this_q = work;
}
DataInfo info{kq, this_q, (size_t)k->ne[1], (size_t)row_size_q, 0, 1, nullptr, 0};
mm.mul_mat_NxM(k->ne[0], (const char *)k->data, k->nb[1], info, k->ne[1], q->ne[1]);
auto kq_i = kq;
int n_step = (k->ne[1] + k_indexer_chunks - 1)/k_indexer_chunks;
for (int i_step = 0; i_step < n_step; ++i_step) {
int nk = std::min(k_indexer_chunks, int(k->ne[1]) - i_step*k_indexer_chunks);
DataInfo info{kq, this_q, (size_t)nk, (size_t)row_size_q, 0, 1, nullptr, 0};
mm.mul_mat_NxM(k->ne[0], (const char *)k->data + i_step*k_indexer_chunks*k->nb[1], k->nb[1], info, nk, q->ne[1]);
if (m->type == GGML_TYPE_F32) {
std::memcpy(score, this_m, k->ne[1]*sizeof(float));
} else {
iqk_f16_to_f32(k->ne[1], (const ggml_fp16_t *)this_m, score);
}
for (int i = 0; i < int(q->ne[1]); ++i) {
float wi = this_w[i];
for (int j = 0; j < int(k->ne[1]); ++j) {
float relu = kq_i[j] > 0.0f ? kq_i[j] : 0.0f;
score[j] += wi * relu;
auto kq_i = kq;
auto this_score = score + i_step*k_indexer_chunks;
if (m->type == GGML_TYPE_F32) {
std::memcpy(this_score, (const float *)this_m + i_step*k_indexer_chunks, nk*sizeof(float));
} else {
iqk_f16_to_f32(nk, (const ggml_fp16_t *)this_m + i_step*k_indexer_chunks, this_score);
}
#ifdef __AVX2__
if constexpr (k_indexer_chunks == 64) {
if (nk == k_indexer_chunks) {
__m256 acc[k_indexer_chunks/8];
for (int j = 0; j < k_indexer_chunks/8; ++j) acc[j] = _mm256_loadu_ps(this_score + 8*j);
for (int i = 0; i < int(q->ne[1]); ++i) {
auto wi = _mm256_set1_ps(this_w[i]);
for (int j = 0; j < k_indexer_chunks/8; ++j) {
auto relu = _mm256_loadu_ps(kq_i + 8*j);
auto mask = _mm256_cmp_ps(relu, _mm256_setzero_ps(), _CMP_GT_OQ);
relu = _mm256_blendv_ps(_mm256_setzero_ps(), relu, mask);
acc[j] = _mm256_fmadd_ps(wi, relu, acc[j]);
}
kq_i += k_indexer_chunks;
}
for (int j = 0; j < k_indexer_chunks/8; ++j) _mm256_storeu_ps(this_score + 8*j, acc[j]);
continue;
}
}
#endif
for (int i = 0; i < int(q->ne[1]); ++i) {
float wi = this_w[i];
for (int j = 0; j < nk; ++j) {
float relu = kq_i[j] > 0.0f ? kq_i[j] : 0.0f;
this_score[j] += wi * relu;
}
kq_i += nk;
}
kq_i += k->ne[1];
}
for (int j = 0; j < int(k->ne[1]); ++j) sorted[j] = j;
std::partial_sort(sorted, sorted + n_top_k, sorted + k->ne[1], [score] (int32_t l, int32_t r) -> bool { return score[l] > score[r]; });
@@ -1903,24 +1930,26 @@ bool iqk_indexer_topk(struct ggml_tensor * dst, void * work_buffer, barrier_t ba
auto score_th = score + first;
auto sorted = (int32_t *)(score + k->ne[1]);
for (int iq = 0; iq < q->ne[2]; ++iq) {
auto this_q = q_data + iq*qnb2;
auto this_m = (const char *)m->data + iq*m->nb[1];
auto this_w = (const float *)((const char *)w->data + w->nb[1]*iq);
DataInfo info{kq_th, this_q, (size_t)n_this_thread, (size_t)row_size_q, 0, 1, nullptr, 0};
mm.mul_mat_NxM(k->ne[0], (const char *)k->data + first*k->nb[1], k->nb[1], info, n_this_thread, q->ne[1]);
if (m->type == GGML_TYPE_F32) {
std::memcpy(score_th, this_m + first*sizeof(float), n_this_thread*sizeof(float));
} else {
iqk_f16_to_f32(n_this_thread, (const ggml_fp16_t *)this_m + first, score_th);
}
auto kq_i = kq_th;
for (int i = 0; i < int(q->ne[1]); ++i) {
float wi = this_w[i];
for (int j = 0; j < n_this_thread; ++j) {
float relu = kq_i[j] > 0.0f ? kq_i[j] : 0.0f;
score_th[j] += wi * relu;
if (n_this_thread > 0) {
auto this_q = q_data + iq*qnb2;
auto this_m = (const char *)m->data + iq*m->nb[1];
auto this_w = (const float *)((const char *)w->data + w->nb[1]*iq);
DataInfo info{kq_th, this_q, (size_t)n_this_thread, (size_t)row_size_q, 0, 1, nullptr, 0};
mm.mul_mat_NxM(k->ne[0], (const char *)k->data + first*k->nb[1], k->nb[1], info, n_this_thread, q->ne[1]);
if (m->type == GGML_TYPE_F32) {
std::memcpy(score_th, this_m + first*sizeof(float), n_this_thread*sizeof(float));
} else {
iqk_f16_to_f32(n_this_thread, (const ggml_fp16_t *)this_m + first, score_th);
}
auto kq_i = kq_th;
for (int i = 0; i < int(q->ne[1]); ++i) {
float wi = this_w[i];
for (int j = 0; j < n_this_thread; ++j) {
float relu = kq_i[j] > 0.0f ? kq_i[j] : 0.0f;
score_th[j] += wi * relu;
}
kq_i += n_this_thread;
}
kq_i += n_this_thread;
}
barrier(barrier_data);
if (ith == 0) {