Faster imatrix (#1801)

* Faster imatrix on AVX2

* Slightly better
This commit is contained in:
Kawrakow 2026-05-15 07:15:16 +03:00 committed by GitHub
parent 35fbe08d6e
commit ba72890076
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
1 changed files with 95 additions and 14 deletions

View File

@ -26,6 +26,79 @@
#pragma warning(disable: 4244 4267) // possible loss of data
#endif
#if defined(__GNUC__) && (defined(__x86_64__) || defined(__i386__))
#include <immintrin.h>
__attribute__((target("avx2,fma")))
static inline bool add_and_check_nans_avx2(int n, const float * x, float * y, int * counts) {
int i = 0;
auto has_nans = _mm256_setzero_ps();
auto one = _mm256_set1_epi32(1);
{
__m256 vx[4], vy[4];
__m256i cy[4];
for ( ; i + 32 < n; i += 32) {
for (int k = 0; k < 4; ++k) {
vx[k] = _mm256_loadu_ps(x + i + 8*k);
vy[k] = _mm256_loadu_ps(y + i + 8*k);
cy[k] = _mm256_loadu_si256((const __m256i *)(counts + i + 8*k));
vy[k] = _mm256_fmadd_ps(vx[k], vx[k], vy[k]);
cy[k] = _mm256_add_epi32(cy[k], one);
auto mask = _mm256_cmp_ps(vx[k], vx[k], _CMP_UNORD_Q);
has_nans = _mm256_or_ps(has_nans, mask);
}
for (int k = 0; k < 4; ++k) {
_mm256_storeu_ps(y + i + 8*k, vy[k]);
_mm256_storeu_si256((__m256i *)(counts + i + 8*k), cy[k]);
}
}
}
for ( ; i + 7 < n; i += 8) {
auto vx = _mm256_loadu_ps(x + i);
auto vy = _mm256_loadu_ps(y + i);
auto cy = _mm256_loadu_si256((const __m256i *)(counts + i));
vy = _mm256_fmadd_ps(vx, vx, vy);
cy = _mm256_add_epi32(cy, one);
_mm256_storeu_ps(y + i, vy);
_mm256_storeu_si256((__m256i *)(counts + i), cy);
auto mask = _mm256_cmp_ps(vx, vx, _CMP_UNORD_Q);
has_nans = _mm256_or_ps(has_nans, mask);
}
auto has_any = _mm256_movemask_ps(has_nans);
if (has_any) {
return true;
}
for (; i < n; ++i) {
if (std::isnan(x[i])) {
return true;
}
y[i] += x[i]*x[i];
++counts[i];
}
return false;
}
#endif
static inline bool add_and_check_nans_scalar(int n, const float * x, float * y, int * counts) {
for (int i = 0; i < n; ++i) {
if (std::isnan(x[i])) {
return true;
}
y[i] += x[i]*x[i];
++counts[i];
}
return false;
}
static bool add_and_check_nans(int n, const float * x, float * y, int * counts) {
#if defined(__GNUC__) && (defined(__x86_64__) || defined(__i386__))
static const bool has_avx2 = __builtin_cpu_supports("avx2");
static const bool has_fma = __builtin_cpu_supports("fma");
if (has_avx2 && has_fma) {
return add_and_check_nans_avx2(n, x, y, counts);
}
#endif
return add_and_check_nans_scalar(n, x, y, counts);
}
uint32_t llama_mtp_state_n_embd(const struct llama_context * ctx);
void llama_set_mtp_target_context(struct llama_context * ctx, struct llama_context * target_ctx);
@ -401,14 +474,18 @@ bool IMatrixCollector::collect_imatrix(struct ggml_tensor * t, bool ask, void *
const int64_t i12 = row;
const float * x = (const float *)((const char *)data + i11*src1->nb[1] + i12*src1->nb[2]);
for (int j = 0; j < (int)src1->ne[0]; ++j) {
e.values[e_start + j] += x[j]*x[j];
e.counts[e_start + j]++;
if (!std::isfinite(e.values[e_start + j])) {
fprintf(stderr, "%f detected in %s\n", e.values[e_start + j], wname.c_str());
exit(1);
}
if (add_and_check_nans(src1->ne[0], x, e.values.data() + e_start, e.counts.data() + e_start)) {
fprintf(stderr, "etected NaNs in %s\n", wname.c_str());
exit(1);
}
//for (int j = 0; j < (int)src1->ne[0]; ++j) {
// e.values[e_start + j] += x[j]*x[j];
// e.counts[e_start + j]++;
// if (!std::isfinite(e.values[e_start + j])) {
// fprintf(stderr, "%f detected in %s\n", e.values[e_start + j], wname.c_str());
// exit(1);
// }
//}
}
}
if (e.ncall > m_last_call) {
@ -482,14 +559,18 @@ bool IMatrixCollector::collect_imatrix(struct ggml_tensor * t, bool ask, void *
auto counts = e.counts.data() + i02*src0->ne[0];
for (int i11 = 0; i11 < (int)src1->ne[1]; ++i11) {
const float * x = (const float *)((const char *)data + i11*src1->nb[1] + i12*src1->nb[2]);
for (int j = 0; j < (int)src1->ne[0]; ++j) {
values[j] += x[j]*x[j];
counts[j]++;
if (!std::isfinite(values[j])) {
fprintf(stderr, "%f detected in %s\n", e.values[j], wname.c_str());
exit(1);
}
if (add_and_check_nans(src1->ne[0], x, values, counts)) {
fprintf(stderr, "detected NaNs in %s\n", wname.c_str());
exit(1);
}
//for (int j = 0; j < (int)src1->ne[0]; ++j) {
// values[j] += x[j]*x[j];
// counts[j]++;
// if (!std::isfinite(values[j])) {
// fprintf(stderr, "%f detected in %s\n", values[j], wname.c_str());
// exit(1);
// }
//}
}
}
if (e.ncall > m_last_call) {