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 }));