Commit 46baf1f1f for llama.cpp

commit 46baf1f1fec5a06d1e52122a9207ca978b720f95
Author: bri-prism <288398250+bri-prism@users.noreply.github.com>
Date:   Wed Oct 7 23:50:22 2026 -0700

    sycl: FWHT optimizations (#29605)

diff --git a/ggml/src/ggml-sycl/fwht.cpp b/ggml/src/ggml-sycl/fwht.cpp
index fb48d7fec..5af549c23 100644
--- a/ggml/src/ggml-sycl/fwht.cpp
+++ b/ggml/src/ggml-sycl/fwht.cpp
@@ -46,8 +46,8 @@ static constexpr float H20[20][20] = {
 #undef P
 #undef N

-template <int N>
-static void fwht_kernel(const float * __restrict__ src, float * __restrict__ dst, const int64_t n_rows,
+template <int N, typename T>
+static void fwht_kernel(const T * __restrict__ src, float * __restrict__ dst, const int64_t n_rows,
                         const float scale, const sycl::nd_item<2> & item) {
     const sycl::sub_group sg = item.get_sub_group();

@@ -67,7 +67,7 @@ static void fwht_kernel(const float * __restrict__ src, float * __restrict__ dst

 #pragma unroll
     for (int i = 0; i < el_w; ++i) {
-        reg[i] = src[i * WARP_SIZE + lane] * scale;
+        reg[i] = static_cast<float>(src[i * WARP_SIZE + lane]) * scale;
     }

     // Butterflies inside the sub-group. The partner of a lane with bit h clear is the
@@ -107,8 +107,8 @@ static void fwht_kernel(const float * __restrict__ src, float * __restrict__ dst
     }
 }

-template <int N>
-static void launch_fwht(const float * src, float * dst, const int64_t n_rows, const float scale,
+template <int N, typename T>
+static void launch_fwht(const T * src, float * dst, const int64_t n_rows, const float scale,
                         dpct::queue_ptr stream) {
     constexpr int rows_per_block = 4;

@@ -120,7 +120,7 @@ static void launch_fwht(const float * src, float * dst, const int64_t n_rows, co

     stream->parallel_for(sycl::nd_range<2>(global, local),
                          [=](sycl::nd_item<2> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
-                             fwht_kernel<N>(src, dst, n_rows, scale, item);
+                             fwht_kernel<N, T>(src, dst, n_rows, scale, item);
                          });
 }

@@ -128,8 +128,8 @@ static void launch_fwht(const float * src, float * dst, const int64_t n_rows, co
 // keeps N/NT values rather than N/WARP_SIZE. Butterflies below the sub-group width
 // still shuffle; those up to NT go through work-group local memory; the rest stay
 // in registers.
-template <int N, int NT>
-static void fwht_kernel_wide(const float * __restrict__ src,
+template <int N, int NT, typename T>
+static void fwht_kernel_wide(const T * __restrict__ src,
                              float * __restrict__ dst,
                              const int64_t            n_rows,
                              const float              scale,
@@ -151,7 +151,7 @@ static void fwht_kernel_wide(const float * __restrict__ src,
     float reg[el_w];
 #pragma unroll
     for (int i = 0; i < el_w; ++i) {
-        reg[i] = src[i * NT + tid] * scale;
+        reg[i] = static_cast<float>(src[i * NT + tid]) * scale;
     }

     const sycl::sub_group sg   = item.get_sub_group();
@@ -207,8 +207,8 @@ static void fwht_kernel_wide(const float * __restrict__ src,
     }
 }

-template <int N, int NT>
-static void launch_fwht_wide(const float *   src,
+template <int N, int NT, typename T>
+static void launch_fwht_wide(const T *       src,
                              float *         dst,
                              const int64_t   n_rows,
                              const float     scale,
@@ -220,13 +220,13 @@ static void launch_fwht_wide(const float *   src,
         sycl::local_accessor<float, 1> smem(sycl::range<1>(N), cgh);
         cgh.parallel_for(sycl::nd_range<2>(global, local),
                          [=](sycl::nd_item<2> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
-                             fwht_kernel_wide<N, NT>(src, dst, n_rows, scale, item, get_pointer(smem));
+                             fwht_kernel_wide<N, NT, T>(src, dst, n_rows, scale, item, get_pointer(smem));
                          });
     });
 }

-template <int N, int m>
-static void kronecker_kernel(const float * __restrict__ src,
+template <int N, int m, typename T>
+static void kronecker_kernel(const T * __restrict__ src,
                              float * __restrict__ dst,
                              const int64_t            n_rows,
                              const float              scale,
@@ -255,7 +255,7 @@ static void kronecker_kernel(const float * __restrict__ src,

 #pragma unroll
         for (int j = 0; j < m; ++j) {
-            reg[i * m + j] = src[b_idx * m + j] * scale;
+            reg[i * m + j] = static_cast<float>(src[b_idx * m + j]) * scale;
         }
     }

@@ -321,8 +321,8 @@ static void kronecker_kernel(const float * __restrict__ src,
     }
 }

-template <int N, int m>
-static void launch_kronecker(const float *   src,
+template <int N, int m, typename T>
+static void launch_kronecker(const T *       src,
                              float *         dst,
                              const int64_t   n_rows,
                              const float     scale,
@@ -337,25 +337,16 @@ static void launch_kronecker(const float *   src,

     stream->parallel_for(sycl::nd_range<2>(global, local),
                          [=](sycl::nd_item<2> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
-                             kronecker_kernel<N, m>(src, dst, n_rows, scale, item);
+                             kronecker_kernel<N, m, T>(src, dst, n_rows, scale, item);
                          });
 }

-bool ggml_sycl_op_fwht(ggml_backend_sycl_context & ctx, const ggml_tensor * src, ggml_tensor * dst) {
-    if (src->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_F32) {
-        return false;
-    }
-    if (!ggml_are_same_shape(src, dst)) {
-        return false;
-    }
-    if (!ggml_is_contiguous(src) || !ggml_is_contiguous(dst)) {
-        return false;
-    }
-
+template <typename T>
+static bool ggml_sycl_op_fwht_impl(ggml_backend_sycl_context & ctx, const ggml_tensor * src, ggml_tensor * dst) {
     const int     n    = (int) src->ne[0];
     const int64_t rows = ggml_nrows(src);

-    const float *   src_d  = (const float *) src->data;
+    const T *       src_d  = (const T *) src->data;
     float *         dst_d  = (float *) dst->data;
     dpct::queue_ptr stream = ctx.stream();

@@ -402,3 +393,24 @@ bool ggml_sycl_op_fwht(ggml_backend_sycl_context & ctx, const ggml_tensor * src,
             return false;
     }
 }
+
+bool ggml_sycl_op_fwht(ggml_backend_sycl_context & ctx, const ggml_tensor * src, ggml_tensor * dst) {
+    if (dst->type != GGML_TYPE_F32) {
+        return false;
+    }
+    if (!ggml_are_same_shape(src, dst)) {
+        return false;
+    }
+    if (!ggml_is_contiguous(src) || !ggml_is_contiguous(dst)) {
+        return false;
+    }
+
+    switch (src->type) {
+        case GGML_TYPE_F32:
+            return ggml_sycl_op_fwht_impl<float>(ctx, src, dst);
+        case GGML_TYPE_F16:
+            return ggml_sycl_op_fwht_impl<sycl::half>(ctx, src, dst);
+        default:
+            return false;
+    }
+}
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index a38766cba..453fb78d7 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -11923,6 +11923,17 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_perf() {
     test_cases.emplace_back(new test_mul_mat_hadamard(GGML_TYPE_F32, GGML_TYPE_F32, 128, 2048, 128));
     test_cases.emplace_back(new test_mul_mat_hadamard(GGML_TYPE_F32, GGML_TYPE_F32, 256, 2048, 256));
     test_cases.emplace_back(new test_mul_mat_hadamard(GGML_TYPE_F32, GGML_TYPE_F32, 512, 2048, 512));
+    test_cases.emplace_back(new test_mul_mat_hadamard(GGML_TYPE_F32, GGML_TYPE_F16, 128, 1, 128));
+    test_cases.emplace_back(new test_mul_mat_hadamard(GGML_TYPE_F32, GGML_TYPE_F16, 64, 1, 64));
+    test_cases.emplace_back(new test_mul_mat_hadamard(GGML_TYPE_F32, GGML_TYPE_F16, 256, 1, 256));
+    test_cases.emplace_back(new test_mul_mat_hadamard(GGML_TYPE_F32, GGML_TYPE_F16, 128, 32, 128));
+    test_cases.emplace_back(new test_mul_mat_hadamard(GGML_TYPE_F32, GGML_TYPE_F16, 64, 2048, 64));
+    test_cases.emplace_back(new test_mul_mat_hadamard(GGML_TYPE_F32, GGML_TYPE_F16, 128, 2048, 128));
+    test_cases.emplace_back(new test_mul_mat_hadamard(GGML_TYPE_F32, GGML_TYPE_F16, 256, 2048, 256));
+    test_cases.emplace_back(new test_mul_mat_hadamard(GGML_TYPE_F32, GGML_TYPE_F16, 512, 2048, 512));
+    test_cases.emplace_back(new test_mul_mat_hadamard(GGML_TYPE_F32, GGML_TYPE_F16, 1024, 2048, 1024));
+    test_cases.emplace_back(new test_mul_mat_hadamard(GGML_TYPE_F32, GGML_TYPE_F16, 4096, 2048, 4096));
+    test_cases.emplace_back(new test_mul_mat_hadamard(GGML_TYPE_F32, GGML_TYPE_F16, 8192, 2048, 8192));

     test_cases.emplace_back(new test_solve_tri(GGML_TYPE_F32, { 64, 64, 4, 4 }, { 32, 64, 4, 4 }));
     test_cases.emplace_back(new test_solve_tri(GGML_TYPE_F32, { 128, 128, 4, 2 }, { 32, 128, 4, 2 }));