diff --git a/ggml/src/ggml-cpu/simd-mappings.h b/ggml/src/ggml-cpu/simd-mappings.h index 10ce4bfc59..89a5afa9ca 100644 --- a/ggml/src/ggml-cpu/simd-mappings.h +++ b/ggml/src/ggml-cpu/simd-mappings.h @@ -1314,6 +1314,33 @@ static inline void __lzs_f16cx4_store(ggml_fp16_t * x, float32x4_t v_y) { #define GGML_F16_ARR (GGML_F16_STEP/GGML_F16_EPR) #endif +// GGML_F16_DOT_* +// like GGML_F16_* but for dot products which need F32 accumulation on AVX512-FP16 + +#if defined(__AVX512FP16__) + +#define GGML_F16_DOT_STEP GGML_F32_STEP +#define GGML_F16_DOT_EPR GGML_F32_EPR +#define GGML_F16_DOT_ARR GGML_F32_ARR +#define GGML_F16_DOT_VEC GGML_F32x16 +#define GGML_F16_DOT_VEC_ZERO GGML_F32x16_ZERO +#define GGML_F16_DOT_VEC_LOAD(p, i) _mm512_cvtph_ps(_mm256_loadu_si256((const __m256i *)(p))) +#define GGML_F16_DOT_VEC_FMA GGML_F32x16_FMA +#define GGML_F16_DOT_VEC_REDUCE GGML_F32x16_REDUCE + +#else + +#define GGML_F16_DOT_STEP GGML_F16_STEP +#define GGML_F16_DOT_EPR GGML_F16_EPR +#define GGML_F16_DOT_ARR GGML_F16_ARR +#define GGML_F16_DOT_VEC GGML_F16_VEC +#define GGML_F16_DOT_VEC_ZERO GGML_F16_VEC_ZERO +#define GGML_F16_DOT_VEC_LOAD GGML_F16_VEC_LOAD +#define GGML_F16_DOT_VEC_FMA GGML_F16_VEC_FMA +#define GGML_F16_DOT_VEC_REDUCE GGML_F16_VEC_REDUCE + +#endif // defined(__AVX512FP16__) + #ifdef __cplusplus } #endif diff --git a/ggml/src/ggml-cpu/vec.cpp b/ggml/src/ggml-cpu/vec.cpp index ff2b636df8..3918a6f4d9 100644 --- a/ggml/src/ggml-cpu/vec.cpp +++ b/ggml/src/ggml-cpu/vec.cpp @@ -342,24 +342,24 @@ void ggml_vec_dot_f16(int n, float * GGML_RESTRICT s, size_t bs, ggml_fp16_t * G } #endif // __riscv_zvfh #else - const int np = (n & ~(GGML_F16_STEP - 1)); + const int np = (n & ~(GGML_F16_DOT_STEP - 1)); - GGML_F16_VEC sum[GGML_F16_ARR] = { GGML_F16_VEC_ZERO }; + GGML_F16_DOT_VEC sum[GGML_F16_DOT_ARR] = { GGML_F16_DOT_VEC_ZERO }; - GGML_F16_VEC ax[GGML_F16_ARR]; - GGML_F16_VEC ay[GGML_F16_ARR]; + GGML_F16_DOT_VEC ax[GGML_F16_DOT_ARR]; + GGML_F16_DOT_VEC ay[GGML_F16_DOT_ARR]; - for (int i = 0; i < np; i += GGML_F16_STEP) { - for (int j = 0; j < GGML_F16_ARR; j++) { - ax[j] = GGML_F16_VEC_LOAD(x + i + j*GGML_F16_EPR, j); - ay[j] = GGML_F16_VEC_LOAD(y + i + j*GGML_F16_EPR, j); + for (int i = 0; i < np; i += GGML_F16_DOT_STEP) { + for (int j = 0; j < GGML_F16_DOT_ARR; j++) { + ax[j] = GGML_F16_DOT_VEC_LOAD(x + i + j*GGML_F16_DOT_EPR, j); + ay[j] = GGML_F16_DOT_VEC_LOAD(y + i + j*GGML_F16_DOT_EPR, j); - sum[j] = GGML_F16_VEC_FMA(sum[j], ax[j], ay[j]); + sum[j] = GGML_F16_DOT_VEC_FMA(sum[j], ax[j], ay[j]); } } // reduce sum0..sum3 to sum0 - GGML_F16_VEC_REDUCE(sumf, sum); + GGML_F16_DOT_VEC_REDUCE(sumf, sum); // leftovers for (int i = np; i < n; ++i) { diff --git a/ggml/src/ggml-cpu/vec.h b/ggml/src/ggml-cpu/vec.h index 5de9cb5b7e..ec1f0a14fa 100644 --- a/ggml/src/ggml-cpu/vec.h +++ b/ggml/src/ggml-cpu/vec.h @@ -276,28 +276,28 @@ inline static void ggml_vec_dot_f16_unroll(const int n, const int xs, float * GG const int np = 0; #endif #else - const int np = (n & ~(GGML_F16_STEP - 1)); + const int np = (n & ~(GGML_F16_DOT_STEP - 1)); - GGML_F16_VEC sum[GGML_VEC_DOT_UNROLL][GGML_F16_ARR] = { { GGML_F16_VEC_ZERO } }; + GGML_F16_DOT_VEC sum[GGML_VEC_DOT_UNROLL][GGML_F16_DOT_ARR] = { { GGML_F16_DOT_VEC_ZERO } }; - GGML_F16_VEC ax[GGML_F16_ARR]; - GGML_F16_VEC ay[GGML_F16_ARR]; + GGML_F16_DOT_VEC ax[GGML_F16_DOT_ARR]; + GGML_F16_DOT_VEC ay[GGML_F16_DOT_ARR]; - for (int i = 0; i < np; i += GGML_F16_STEP) { - for (int j = 0; j < GGML_F16_ARR; j++) { - ay[j] = GGML_F16_VEC_LOAD(y + i + j*GGML_F16_EPR, j); + for (int i = 0; i < np; i += GGML_F16_DOT_STEP) { + for (int j = 0; j < GGML_F16_DOT_ARR; j++) { + ay[j] = GGML_F16_DOT_VEC_LOAD(y + i + j*GGML_F16_DOT_EPR, j); for (int k = 0; k < GGML_VEC_DOT_UNROLL; ++k) { - ax[j] = GGML_F16_VEC_LOAD(x[k] + i + j*GGML_F16_EPR, j); + ax[j] = GGML_F16_DOT_VEC_LOAD(x[k] + i + j*GGML_F16_DOT_EPR, j); - sum[k][j] = GGML_F16_VEC_FMA(sum[k][j], ax[j], ay[j]); + sum[k][j] = GGML_F16_DOT_VEC_FMA(sum[k][j], ax[j], ay[j]); } } } // reduce sum0..sum3 to sum0 for (int k = 0; k < GGML_VEC_DOT_UNROLL; ++k) { - GGML_F16_VEC_REDUCE(sumf[k], sum[k]); + GGML_F16_DOT_VEC_REDUCE(sumf[k], sum[k]); } #endif #else