diff --git a/ggml/src/iqk/iqk_mul_mat.cpp b/ggml/src/iqk/iqk_mul_mat.cpp index bbfba28ca..ad5fb12f8 100644 --- a/ggml/src/iqk/iqk_mul_mat.cpp +++ b/ggml/src/iqk/iqk_mul_mat.cpp @@ -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) {