]> git.djapps.eu Git - pkg/ggml/sources/llama.cpp/commitdiff
sycl : enhance OP set_rows to support all missed data types (#26515)
authorNeo Zhang <redacted>
Fri, 7 Aug 2026 04:52:52 +0000 (12:52 +0800)
committerGitHub <redacted>
Fri, 7 Aug 2026 04:52:52 +0000 (07:52 +0300)
* 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
ggml/src/ggml-sycl/set_rows.cpp

index d91e41f9575a26c81787a997ecfc4a304a5dec36..ce92d438d3d992e0185d39b398b35af82f6dea46 100644 (file)
@@ -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;
index 5fb97790712d789ff8a0cfb6576714d58fc43fd7..52a0bcb6ebad870f3038ee57a37976ce45693c80 100644 (file)
@@ -1,6 +1,10 @@
 #include "set_rows.hpp"
 #include "cpy.hpp"
 
+#include "ggml-quants.h"
+
+#include <vector>
+
 namespace utils {
 template<typename T>
 static constexpr bool is_arithmetic_v() {
@@ -20,7 +24,17 @@ convert (const char* src, char* dst) {
    *reinterpret_cast<TOut*>(dst) = dst_val;
 }
 
-template <typename TIdx, typename blockType, int qk, cpy_kernel_t cpyblck>
+#ifdef GGML_SYCL_HAS_BF16
+// sycl::vec::convert does not provide a half -> bfloat16 path, so route through float.
+template<>
+inline void convert<sycl::half, sycl::ext::oneapi::bfloat16>(const char* src, char* dst) {
+    const float tmp = sycl::vec<sycl::half, 1>(*reinterpret_cast<const sycl::half*>(src))
+                          .template convert<float, sycl::rounding_mode::automatic>()[0];
+    *reinterpret_cast<sycl::ext::oneapi::bfloat16*>(dst) = sycl::ext::oneapi::bfloat16(tmp);
+}
+#endif
+
+template <typename TIn, typename TIdx, typename blockType, int qk, cpy_kernel_t cpyblck>
 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<char *>(reinterpret_cast<char *>(dst_d) + dst_offset);
-        cpyblck(src_block, dst_block);
+        if constexpr (std::is_same_v<TIn, float>) {
+            cpyblck(src_block, dst_block);
+        } else {
+            float src_block_f32[qk];
+            const TIn * src_block_t = reinterpret_cast<const TIn *>(src_block);
+            for (int j = 0; j < qk; ++j) {
+                src_block_f32[j] = (float) src_block_t[j];
+            }
+            cpyblck(reinterpret_cast<const char *>(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<typename blockType>
+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 <typename TIn, typename TIdx, typename blockType, int qk, quantize_row_qk_t<blockType> 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<char> src0_host(src0_bytes);
+    std::vector<char> 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<float> src_row_f32(ne00);
+    const int64_t nblocks = ne00 / qk;
+    std::vector<blockType> 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<const TIn *>(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 <typename TIn, typename TIdx, typename blockType, int qk, quantize_rows_f_t quantize_rows>
+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<char> src0_host(src0_bytes);
+    std::vector<char> 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<float> src_row_f32(ne00);
+    const int64_t nblocks = ne00 / qk;
+    std::vector<blockType> 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<const TIn *>(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<typename TIn, typename TIdx, typename TOut>
 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<TIdx, block_q8_0, QK8_0, cpy_blck_f32_q8_0>(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<TIn, TIdx, block_q8_0, QK8_0, cpy_blck_f32_q8_0>(
+                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<TIdx, block_q1_0, QK1_0, cpy_blck_f32_q1_0>(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<TIn, TIdx, block_q1_0, QK1_0, cpy_blck_f32_q1_0>(
+                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<TIn, TIdx, block_q2_0, QK2_0, cpy_blck_f32_q2_0>(
+                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<TIdx, block_q5_1, QK5_1, cpy_blck_f32_q5_1>(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<TIn, TIdx, block_q5_1, QK5_1, cpy_blck_f32_q5_1>(
+                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<TIdx, block_q5_0, QK5_0, cpy_blck_f32_q5_0>(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<TIn, TIdx, block_q5_0, QK5_0, cpy_blck_f32_q5_0>(
+                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<TIdx, block_q4_1, QK4_1, cpy_blck_f32_q4_1>(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<TIn, TIdx, block_q4_1, QK4_1, cpy_blck_f32_q4_1>(
+                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<TIdx, block_q4_0, QK4_0, cpy_blck_f32_q4_0>(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<TIn, TIdx, block_q4_0, QK4_0, cpy_blck_f32_q4_0>(
+                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<TIdx, block_iq4_nl, QK4_NL, cpy_blck_f32_iq4_nl>(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<TIn, TIdx, block_iq4_nl, QK4_NL, cpy_blck_f32_iq4_nl>(
+                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<TIdx, block_mxfp4, QK_MXFP4, cpy_blck_f32_mxfp4>(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<TIn, TIdx, block_mxfp4, QK_MXFP4, cpy_blck_f32_mxfp4>(
+                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<TIdx, block_nvfp4, QK_NVFP4, cpy_blck_f32_nvfp4>(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<TIn, TIdx, block_nvfp4, QK_NVFP4, cpy_blck_f32_nvfp4>(
+                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<TIn, TIdx, block_q2_K, QK_K, quantize_row_q2_K_ref>(
+                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<TIn, TIdx, block_q3_K, QK_K, quantize_row_q3_K_ref>(
+                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<TIn, TIdx, block_q4_K, QK_K, quantize_row_q4_K_ref>(
+                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<TIn, TIdx, block_q5_K, QK_K, quantize_row_q5_K_ref>(
+                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<TIn, TIdx, block_q6_K, QK_K, quantize_row_q6_K_ref>(
+                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<TIn, TIdx, block_iq2_xxs, QK_K, quantize_iq2_xxs>(
+                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<TIn, TIdx, block_iq2_xs, QK_K, quantize_iq2_xs>(
+                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<TIn, TIdx, block_iq2_s, QK_K, quantize_iq2_s>(
+                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<TIn, TIdx, block_iq3_xxs, QK_K, quantize_row_iq3_xxs_ref>(
+                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<TIn, TIdx, block_iq3_s, QK_K, quantize_row_iq3_s_ref>(
+                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<TIn, TIdx, block_iq1_s, QK_K, quantize_iq1_s>(
+                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<TIn, TIdx, block_iq1_m, QK_K, quantize_iq1_m>(
+                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<TIn, TIdx, block_iq4_xs, QK_K, quantize_row_iq4_xs_ref>(
+                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<float, int64_t>(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<sycl::half, int64_t>(ctx, src0, src1, dst);
+        } else {
+            set_rows_sycl<sycl::half, int32_t>(ctx, src0, src1, dst);
+        }
     } else {
-        set_rows_sycl<float, int32_t>(ctx, src0, src1, dst);
+        if (src1->type == GGML_TYPE_I64) {
+            set_rows_sycl<float, int64_t>(ctx, src0, src1, dst);
+        } else {
+            set_rows_sycl<float, int32_t>(ctx, src0, src1, dst);
+        }
     }
 }