From f162a194754838c2468d1d02600182154ba41a63 Mon Sep 17 00:00:00 2001 From: Titaniumtown Date: Tue, 1 Sep 2026 09:04:31 -0700 Subject: [PATCH] =?UTF-8?q?Revert=20"sycl=20:=20add=20Kronecker=20product?= =?UTF-8?q?=20FWHT=20support=20for=20sizes=20384,=20640,=20768,=2012?= =?UTF-8?q?=E2=80=A6"=20(#28184)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit This reverts commit 1f3d318734c61cf6f3b209726cdd3f9c300a782e. --- ggml/src/ggml-sycl/fwht.cpp | 172 ------------------------------------ 1 file changed, 172 deletions(-) diff --git a/ggml/src/ggml-sycl/fwht.cpp b/ggml/src/ggml-sycl/fwht.cpp index 39f273bea..2312b3d13 100644 --- a/ggml/src/ggml-sycl/fwht.cpp +++ b/ggml/src/ggml-sycl/fwht.cpp @@ -1,50 +1,6 @@ #include "fwht.hpp" #include -#define P 1.0f -#define N -1.0f - -// constant Hadamard matrix via Paley I construction -static constexpr float H12[12][12] = { - { P, P, P, P, P, P, P, P, P, P, P, P }, - { P, N, P, N, P, P, P, N, N, N, P, N }, - { P, N, N, P, N, P, P, P, N, N, N, P }, - { P, P, N, N, P, N, P, P, P, N, N, N }, - { P, N, P, N, N, P, N, P, P, P, N, N }, - { P, N, N, P, N, N, P, N, P, P, P, N }, - { P, N, N, N, P, N, N, P, N, P, P, P }, - { P, P, N, N, N, P, N, N, P, N, P, P }, - { P, P, P, N, N, N, P, N, N, P, N, P }, - { P, P, P, P, N, N, N, P, N, N, P, N }, - { P, N, P, P, P, N, N, N, P, N, N, P }, - { P, P, N, P, P, P, N, N, N, P, N, N } -}; - -static constexpr float H20[20][20] = { - { P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P, P }, - { P, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N }, - { P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P }, - { P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P }, - { P, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N }, - { P, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N, N }, - { P, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N, N }, - { P, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P, N }, - { P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N, P }, - { P, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P, N }, - { P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N, P }, - { P, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P, N }, - { P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P, P }, - { P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P, P }, - { P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P, P }, - { P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N, P }, - { P, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N, N }, - { P, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P, N }, - { P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N, P }, - { P, P, N, N, P, P, P, P, N, P, N, P, N, N, N, N, P, P, N, N } -}; - -#undef P -#undef N template static void fwht_kernel(const float * __restrict__ src, float * __restrict__ dst, const int64_t n_rows, @@ -124,122 +80,6 @@ static void launch_fwht(const float * src, float * dst, const int64_t n_rows, co }); } -template -static void kronecker_kernel(const float * __restrict__ src, - float * __restrict__ dst, - const int64_t n_rows, - const float scale, - const sycl::nd_item<2> & item) { - static_assert(m == 12 || m == 20, "block size has to be 12 or 20."); - - const sycl::sub_group sg = item.get_sub_group(); - - const int64_t r = item.get_global_id(0); - if (r >= n_rows) { - return; - } - - src += r * N; - dst += r * N; - - constexpr int blocks_per_group = N / m; - constexpr int el_w = blocks_per_group / WARP_SIZE; - static_assert(el_w >= 1 && blocks_per_group % WARP_SIZE == 0, "blocks_per_group must be a multiple of WARP_SIZE"); - float reg[el_w * m]; - const int lane = sg.get_local_linear_id(); - -#pragma unroll - for (int i = 0; i < el_w; ++i) { - const int b_idx = i * WARP_SIZE + lane; - -#pragma unroll - for (int j = 0; j < m; ++j) { - reg[i * m + j] = src[b_idx * m + j] * scale; - } - } - -#pragma unroll - for (int b = 0; b < el_w; ++b) { - float z[m] = { 0.0f }; - -#pragma unroll - for (int i = 0; i < m; ++i) { -#pragma unroll - for (int j = 0; j < m; ++j) { - const float h = (m == 12 ? H12[j][i] : H20[j][i]); - z[i] += reg[b * m + j] * h; - } - } - -#pragma unroll - for (int i = 0; i < m; ++i) { - reg[b * m + i] = z[i]; - } - } - -#pragma unroll - for (int h = 1; h < WARP_SIZE; h *= 2) { -#pragma unroll - for (int j = 0; j < el_w; ++j) { -#pragma unroll - for (int k = 0; k < m; ++k) { - const float val = reg[j * m + k]; - const float val2 = dpct::permute_sub_group_by_xor(sg, val, h, WARP_SIZE); - - reg[j * m + k] = (lane & h) == 0 ? val + val2 : val2 - val; - } - } - } - -#pragma unroll - for (int h = WARP_SIZE; h < blocks_per_group; h *= 2) { - const int step = h / WARP_SIZE; -#pragma unroll - for (int j = 0; j < el_w; j += 2 * step) { -#pragma unroll - for (int s = 0; s < step; ++s) { -#pragma unroll - for (int k = 0; k < m; ++k) { - const float x = reg[(j + s) * m + k]; - const float y = reg[(j + s + step) * m + k]; - - reg[(j + s) * m + k] = x + y; - reg[(j + s + step) * m + k] = x - y; - } - } - } - } - -#pragma unroll - for (int i = 0; i < el_w; ++i) { - const int b_idx = i * WARP_SIZE + lane; -#pragma unroll - for (int k = 0; k < m; ++k) { - dst[b_idx * m + k] = reg[i * m + k]; - } - } -} - -template -static void launch_kronecker(const float * src, - float * dst, - const int64_t n_rows, - const float scale, - dpct::queue_ptr stream) { - constexpr int rows_per_block = 4; - - const int64_t num_blocks = (n_rows + rows_per_block - 1) / rows_per_block; - - // dim 1 is the fastest-varying, so a sub-group is exactly one row's WARP_SIZE lanes. - const sycl::range<2> global(num_blocks * rows_per_block, WARP_SIZE); - const sycl::range<2> local(rows_per_block, WARP_SIZE); - - stream->parallel_for(sycl::nd_range<2>(global, local), - [=](sycl::nd_item<2> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] { - kronecker_kernel(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; @@ -273,18 +113,6 @@ bool ggml_sycl_op_fwht(ggml_backend_sycl_context & ctx, const ggml_tensor * src, case 512: launch_fwht<512>(src_d, dst_d, rows, scale, stream); return true; - case 384: - launch_kronecker<384, 12>(src_d, dst_d, rows, scale, stream); - return true; - case 768: - launch_kronecker<768, 12>(src_d, dst_d, rows, scale, stream); - return true; - case 640: - launch_kronecker<640, 20>(src_d, dst_d, rows, scale, stream); - return true; - case 1280: - launch_kronecker<1280, 20>(src_d, dst_d, rows, scale, stream); - return true; default: return false; }