From c074cb3f763d93262133f9715dd5ca94285f938d Mon Sep 17 00:00:00 2001 From: Neo Zhang Date: Fri, 7 Aug 2026 12:52:52 +0800 Subject: [PATCH] sycl : enhance OP set_rows to support all missed data types (#26515) * support fp16 to fp16/fp32 * support all missed data types in set_rows * refactor the code to support all data types --- ggml/src/ggml-sycl/ggml-sycl.cpp | 11 +- ggml/src/ggml-sycl/set_rows.cpp | 360 +++++++++++++++++++++++++++++-- 2 files changed, 347 insertions(+), 24 deletions(-) diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp index d91e41f957..ce92d438d3 100644 --- a/ggml/src/ggml-sycl/ggml-sycl.cpp +++ b/ggml/src/ggml-sycl/ggml-sycl.cpp @@ -5795,14 +5795,9 @@ static bool do_ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, cons case GGML_OP_SET_ROWS: { - - auto res = ((op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_F16 || op->type == GGML_TYPE_BF16 || - op->type == GGML_TYPE_Q8_0 || op->type == GGML_TYPE_Q5_1 || op->type == GGML_TYPE_Q5_0 || - op->type == GGML_TYPE_Q1_0 || - op->type == GGML_TYPE_Q4_1 || op->type == GGML_TYPE_Q4_0 || op->type == GGML_TYPE_IQ4_NL || - op->type == GGML_TYPE_MXFP4 || op->type == GGML_TYPE_NVFP4) && - op->src[0]->type == GGML_TYPE_F32 && - (op->src[1]->type == GGML_TYPE_I64 || op->src[1]->type == GGML_TYPE_I32)); + auto res = (op->src[0]->type == GGML_TYPE_F32 || op->src[0]->type == GGML_TYPE_F16 || + op->src[0]->type == GGML_TYPE_BF16) && + (op->src[1]->type == GGML_TYPE_I64 || op->src[1]->type == GGML_TYPE_I32); return res; } break; diff --git a/ggml/src/ggml-sycl/set_rows.cpp b/ggml/src/ggml-sycl/set_rows.cpp index 5fb9779071..52a0bcb6eb 100644 --- a/ggml/src/ggml-sycl/set_rows.cpp +++ b/ggml/src/ggml-sycl/set_rows.cpp @@ -1,6 +1,10 @@ #include "set_rows.hpp" #include "cpy.hpp" +#include "ggml-quants.h" + +#include + namespace utils { template static constexpr bool is_arithmetic_v() { @@ -20,7 +24,17 @@ convert (const char* src, char* dst) { *reinterpret_cast(dst) = dst_val; } -template +#ifdef GGML_SYCL_HAS_BF16 +// sycl::vec::convert does not provide a half -> bfloat16 path, so route through float. +template<> +inline void convert(const char* src, char* dst) { + const float tmp = sycl::vec(*reinterpret_cast(src)) + .template convert()[0]; + *reinterpret_cast(dst) = sycl::ext::oneapi::bfloat16(tmp); +} +#endif + +template static void set_rows_sycl_q(const char * __restrict__ src0_d, const TIdx * __restrict__ src1_d, blockType * __restrict__ dst_d, @@ -68,13 +82,22 @@ static void set_rows_sycl_q(const char * __restrict__ src0_d, const int64_t i11 = i02 % ne11; const int64_t i10 = i01; const size_t src_offset = calculate_offset<3>({ nb01, nb02, nb03 }, { i01, i02, i03 }); - const char * src_block = src0_d + src_offset + i00 * sizeof(float); + const char * src_block = src0_d + src_offset + i00 * sizeof(TIn); const size_t src1_offset = calculate_offset<3>({ nb10, nb11, nb12 }, { i10, i11, i12 }); const int64_t dst_row = src1_d[src1_offset / sizeof(TIdx)]; const size_t dst_offset = calculate_offset<3>({ nb1, nb2, nb3 }, { dst_row, i02, i03 }) + (i00 / qk) * sizeof(blockType); char * dst_block = reinterpret_cast(reinterpret_cast(dst_d) + dst_offset); - cpyblck(src_block, dst_block); + if constexpr (std::is_same_v) { + cpyblck(src_block, dst_block); + } else { + float src_block_f32[qk]; + const TIn * src_block_t = reinterpret_cast(src_block); + for (int j = 0; j < qk; ++j) { + src_block_f32[j] = (float) src_block_t[j]; + } + cpyblck(reinterpret_cast(src_block_f32), dst_block); + } }); GGML_UNUSED(ne10); GGML_UNUSED(ne13); @@ -82,6 +105,139 @@ static void set_rows_sycl_q(const char * __restrict__ src0_d, GGML_UNUSED(nb13); } +template +using quantize_row_qk_t = void (*)(const float *, blockType *, int64_t); + +using quantize_rows_f_t = size_t (*)(const float *, void *, int64_t, int64_t, const float *); + +template quantize_row> +static void set_rows_sycl_qk_host( + const ggml_tensor * src0, + const ggml_tensor * src1, + ggml_tensor * dst, + const int64_t ne00, + const int64_t ne01, + const int64_t ne02, + const int64_t ne03, + const int64_t ne11, + const int64_t ne12, + const size_t nb01, + const size_t nb02, + const size_t nb03, + const size_t nb10, + const size_t nb11, + const size_t nb12, + const size_t nb1, + const size_t nb2, + const size_t nb3, + queue_ptr stream) { + GGML_ASSERT(ne00 % qk == 0); + + const size_t src0_bytes = ggml_nbytes(src0); + const size_t src1_bytes = ggml_nbytes(src1); + + std::vector src0_host(src0_bytes); + std::vector src1_host(src1_bytes); + + stream->memcpy(src0_host.data(), src0->data, src0_bytes); + stream->memcpy(src1_host.data(), src1->data, src1_bytes); + stream->wait(); + + std::vector src_row_f32(ne00); + const int64_t nblocks = ne00 / qk; + std::vector dst_row_q(nblocks); + + for (int64_t i03 = 0; i03 < ne03; ++i03) { + for (int64_t i02 = 0; i02 < ne02; ++i02) { + for (int64_t i01 = 0; i01 < ne01; ++i01) { + const int64_t i12 = i03 % ne12; + const int64_t i11 = i02 % ne11; + const int64_t i10 = i01; + + const size_t src1_offset = calculate_offset<3>({ nb10, nb11, nb12 }, { i10, i11, i12 }); + const int64_t dst_row = *(const TIdx *) (src1_host.data() + src1_offset); + + const size_t src0_row_offset = calculate_offset<3>({ nb01, nb02, nb03 }, { i01, i02, i03 }); + const TIn * src_row = reinterpret_cast(src0_host.data() + src0_row_offset); + + for (int64_t i00 = 0; i00 < ne00; ++i00) { + src_row_f32[i00] = (float) src_row[i00]; + } + + quantize_row(src_row_f32.data(), dst_row_q.data(), ne00); + + const size_t dst_offset = calculate_offset<3>({ nb1, nb2, nb3 }, { dst_row, i02, i03 }); + stream->memcpy((char *) dst->data + dst_offset, dst_row_q.data(), nblocks * sizeof(blockType)); + stream->wait(); + } + } + } +} + +template +static void set_rows_sycl_iq_host( + const ggml_tensor * src0, + const ggml_tensor * src1, + ggml_tensor * dst, + const int64_t ne00, + const int64_t ne01, + const int64_t ne02, + const int64_t ne03, + const int64_t ne11, + const int64_t ne12, + const size_t nb01, + const size_t nb02, + const size_t nb03, + const size_t nb10, + const size_t nb11, + const size_t nb12, + const size_t nb1, + const size_t nb2, + const size_t nb3, + queue_ptr stream) { + GGML_ASSERT(ne00 % qk == 0); + + const size_t src0_bytes = ggml_nbytes(src0); + const size_t src1_bytes = ggml_nbytes(src1); + + std::vector src0_host(src0_bytes); + std::vector src1_host(src1_bytes); + + stream->memcpy(src0_host.data(), src0->data, src0_bytes); + stream->memcpy(src1_host.data(), src1->data, src1_bytes); + stream->wait(); + + std::vector src_row_f32(ne00); + const int64_t nblocks = ne00 / qk; + std::vector dst_row_q(nblocks); + + for (int64_t i03 = 0; i03 < ne03; ++i03) { + for (int64_t i02 = 0; i02 < ne02; ++i02) { + for (int64_t i01 = 0; i01 < ne01; ++i01) { + const int64_t i12 = i03 % ne12; + const int64_t i11 = i02 % ne11; + const int64_t i10 = i01; + + const size_t src1_offset = calculate_offset<3>({ nb10, nb11, nb12 }, { i10, i11, i12 }); + const int64_t dst_row = *(const TIdx *) (src1_host.data() + src1_offset); + + const size_t src0_row_offset = calculate_offset<3>({ nb01, nb02, nb03 }, { i01, i02, i03 }); + const TIn * src_row = reinterpret_cast(src0_host.data() + src0_row_offset); + + for (int64_t i00 = 0; i00 < ne00; ++i00) { + src_row_f32[i00] = (float) src_row[i00]; + } + + quantize_rows(src_row_f32.data(), dst_row_q.data(), 1, ne00, nullptr); + + const size_t dst_offset = calculate_offset<3>({ nb1, nb2, nb3 }, { dst_row, i02, i03 }); + stream->memcpy((char *) dst->data + dst_offset, dst_row_q.data(), nblocks * sizeof(blockType)); + stream->wait(); + } + } + } +} + template static void k_set_rows( const char * __restrict__ src0, const TIdx * __restrict__ src1, char * __restrict__ dst, @@ -200,31 +356,194 @@ static void set_rows_sycl(ggml_backend_sycl_context & ctx, const ggml_tensor * s break; #endif case GGML_TYPE_Q8_0: - set_rows_sycl_q(src0_d, src1_d, (block_q8_0 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); + set_rows_sycl_q( + src0_d, src1_d, (block_q8_0 *) dst->data, ne00, ne01, ne02, ne03, + ne10, ne11, ne12, ne13, nb00, nb01, + nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); break; case GGML_TYPE_Q1_0: - set_rows_sycl_q(src0_d, src1_d, (block_q1_0 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); + set_rows_sycl_q( + src0_d, src1_d, (block_q1_0 *) dst->data, ne00, ne01, ne02, ne03, + ne10, ne11, ne12, ne13, nb00, nb01, + nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); + break; + case GGML_TYPE_Q2_0: + set_rows_sycl_q( + src0_d, src1_d, (block_q2_0 *) dst->data, ne00, ne01, ne02, ne03, + ne10, ne11, ne12, ne13, nb00, nb01, + nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); break; case GGML_TYPE_Q5_1: - set_rows_sycl_q(src0_d, src1_d, (block_q5_1 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); + set_rows_sycl_q( + src0_d, src1_d, (block_q5_1 *) dst->data, ne00, ne01, ne02, ne03, + ne10, ne11, ne12, ne13, nb00, nb01, + nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); break; case GGML_TYPE_Q5_0: - set_rows_sycl_q(src0_d, src1_d, (block_q5_0 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); + set_rows_sycl_q( + src0_d, src1_d, (block_q5_0 *) dst->data, ne00, ne01, ne02, ne03, + ne10, ne11, ne12, ne13, nb00, nb01, + nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); break; case GGML_TYPE_Q4_1: - set_rows_sycl_q(src0_d, src1_d, (block_q4_1 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); + set_rows_sycl_q( + src0_d, src1_d, (block_q4_1 *) dst->data, ne00, ne01, ne02, ne03, + ne10, ne11, ne12, ne13, nb00, nb01, + nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); break; case GGML_TYPE_Q4_0: - set_rows_sycl_q(src0_d, src1_d, (block_q4_0 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); + set_rows_sycl_q( + src0_d, src1_d, (block_q4_0 *) dst->data, ne00, ne01, ne02, ne03, + ne10, ne11, ne12, ne13, nb00, nb01, + nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); break; case GGML_TYPE_IQ4_NL: - set_rows_sycl_q(src0_d, src1_d, (block_iq4_nl *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); + set_rows_sycl_q( + src0_d, src1_d, (block_iq4_nl *) dst->data, ne00, ne01, ne02, ne03, + ne10, ne11, ne12, ne13, nb00, nb01, + nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); break; case GGML_TYPE_MXFP4: - set_rows_sycl_q(src0_d, src1_d, (block_mxfp4 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); + set_rows_sycl_q( + src0_d, src1_d, (block_mxfp4 *) dst->data, ne00, ne01, ne02, ne03, + ne10, ne11, ne12, ne13, nb00, nb01, + nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); break; case GGML_TYPE_NVFP4: - set_rows_sycl_q(src0_d, src1_d, (block_nvfp4 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); + set_rows_sycl_q( + src0_d, src1_d, (block_nvfp4 *) dst->data, ne00, ne01, ne02, ne03, + ne10, ne11, ne12, ne13, nb00, nb01, + nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream); + break; + case GGML_TYPE_Q2_K: + set_rows_sycl_qk_host( + src0, src1, dst, + ne00, ne01, ne02, ne03, + ne11, ne12, + nb01, nb02, nb03, + nb10, nb11, nb12, + nb1, nb2, nb3, + stream); + break; + case GGML_TYPE_Q3_K: + set_rows_sycl_qk_host( + src0, src1, dst, + ne00, ne01, ne02, ne03, + ne11, ne12, + nb01, nb02, nb03, + nb10, nb11, nb12, + nb1, nb2, nb3, + stream); + break; + case GGML_TYPE_Q4_K: + set_rows_sycl_qk_host( + src0, src1, dst, + ne00, ne01, ne02, ne03, + ne11, ne12, + nb01, nb02, nb03, + nb10, nb11, nb12, + nb1, nb2, nb3, + stream); + break; + case GGML_TYPE_Q5_K: + set_rows_sycl_qk_host( + src0, src1, dst, + ne00, ne01, ne02, ne03, + ne11, ne12, + nb01, nb02, nb03, + nb10, nb11, nb12, + nb1, nb2, nb3, + stream); + break; + case GGML_TYPE_Q6_K: + set_rows_sycl_qk_host( + src0, src1, dst, + ne00, ne01, ne02, ne03, + ne11, ne12, + nb01, nb02, nb03, + nb10, nb11, nb12, + nb1, nb2, nb3, + stream); + break; + case GGML_TYPE_IQ2_XXS: + set_rows_sycl_iq_host( + src0, src1, dst, + ne00, ne01, ne02, ne03, + ne11, ne12, + nb01, nb02, nb03, + nb10, nb11, nb12, + nb1, nb2, nb3, + stream); + break; + case GGML_TYPE_IQ2_XS: + set_rows_sycl_iq_host( + src0, src1, dst, + ne00, ne01, ne02, ne03, + ne11, ne12, + nb01, nb02, nb03, + nb10, nb11, nb12, + nb1, nb2, nb3, + stream); + break; + case GGML_TYPE_IQ2_S: + set_rows_sycl_iq_host( + src0, src1, dst, + ne00, ne01, ne02, ne03, + ne11, ne12, + nb01, nb02, nb03, + nb10, nb11, nb12, + nb1, nb2, nb3, + stream); + break; + case GGML_TYPE_IQ3_XXS: + set_rows_sycl_qk_host( + src0, src1, dst, + ne00, ne01, ne02, ne03, + ne11, ne12, + nb01, nb02, nb03, + nb10, nb11, nb12, + nb1, nb2, nb3, + stream); + break; + case GGML_TYPE_IQ3_S: + set_rows_sycl_qk_host( + src0, src1, dst, + ne00, ne01, ne02, ne03, + ne11, ne12, + nb01, nb02, nb03, + nb10, nb11, nb12, + nb1, nb2, nb3, + stream); + break; + case GGML_TYPE_IQ1_S: + set_rows_sycl_iq_host( + src0, src1, dst, + ne00, ne01, ne02, ne03, + ne11, ne12, + nb01, nb02, nb03, + nb10, nb11, nb12, + nb1, nb2, nb3, + stream); + break; + case GGML_TYPE_IQ1_M: + set_rows_sycl_iq_host( + src0, src1, dst, + ne00, ne01, ne02, ne03, + ne11, ne12, + nb01, nb02, nb03, + nb10, nb11, nb12, + nb1, nb2, nb3, + stream); + break; + case GGML_TYPE_IQ4_XS: + set_rows_sycl_qk_host( + src0, src1, dst, + ne00, ne01, ne02, ne03, + ne11, ne12, + nb01, nb02, nb03, + nb10, nb11, nb12, + nb1, nb2, nb3, + stream); break; default: GGML_ABORT("Unsupported tensor type!"); @@ -237,12 +556,21 @@ void ggml_sycl_op_set_rows(ggml_backend_sycl_context & ctx, ggml_tensor * dst) { const ggml_tensor * src0 = dst->src[0]; const ggml_tensor * src1 = dst->src[1]; - GGML_ASSERT(dst->src[0]->type == GGML_TYPE_F32); + GGML_ASSERT(dst->src[0]->type == GGML_TYPE_F32 || dst->src[0]->type == GGML_TYPE_F16); GGML_ASSERT(dst->src[1]->type == GGML_TYPE_I64 || dst->src[1]->type == GGML_TYPE_I32); - if (src1->type == GGML_TYPE_I64) { - set_rows_sycl(ctx, src0, src1, dst); + // dispatch on the index type (src1) and the source value type (src0) + if (src0->type == GGML_TYPE_F16) { + if (src1->type == GGML_TYPE_I64) { + set_rows_sycl(ctx, src0, src1, dst); + } else { + set_rows_sycl(ctx, src0, src1, dst); + } } else { - set_rows_sycl(ctx, src0, src1, dst); + if (src1->type == GGML_TYPE_I64) { + set_rows_sycl(ctx, src0, src1, dst); + } else { + set_rows_sycl(ctx, src0, src1, dst); + } } }