mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-02 02:47:26 -05:00
ggml : accumulate f16 dot products in f32 on AVX512-FP16 (#29545)
Supersedes #29530 Signed-off-by: Adrien Gallouët <angt@huggingface.co>
This commit is contained in:
@@ -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
|
||||
|
||||
+10
-10
@@ -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) {
|
||||
|
||||
+10
-10
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user