diff --git a/ggml/src/ggml-sycl/common.hpp b/ggml/src/ggml-sycl/common.hpp index 8f2106d83088..a74aef9d975b 100644 --- a/ggml/src/ggml-sycl/common.hpp +++ b/ggml/src/ggml-sycl/common.hpp @@ -224,7 +224,7 @@ inline dpct::err0 ggml_sycl_set_device(const int device) try { ////////////////////// struct optimize_feature { bool reorder=false; - bool xmx_pq2=false; // PQ2_0 rewritten into the XMX layout (pq2_xmx.hpp); only that path can read it + bool xmx_pq2=false; // PQ2_0/PTQ1_0 rewritten into the XMX layout (pq2_xmx.hpp); only that path can read it }; struct sycl_device_info { diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp index 97b5f35659f7..3322ae40824f 100644 --- a/ggml/src/ggml-sycl/ggml-sycl.cpp +++ b/ggml/src/ggml-sycl/ggml-sycl.cpp @@ -983,16 +983,12 @@ static bool ggml_sycl_xmx_disabled() { return disabled; } -// PTQ1_0 weights are expanded into the 34-byte PQ2_0 XMX blocks on first use (pq2_xmx.hpp), so on devices -// that run that path their allocation reserves room for the expanded form -static bool ggml_sycl_ptq1_xmx_expands(int device, const ggml_tensor * tensor) { - return tensor->type == GGML_TYPE_PTQ1_0 && g_ggml_sycl_enable_optimize && !ggml_sycl_xmx_disabled() && - tensor->ne[2] == 1 && tensor->ne[3] == 1 && ggml_sycl_pq2_xmx_supports_ne0(tensor->ne[0]) && - ggml_sycl_device_has_dpas16(device); -} - -static size_t ggml_sycl_ptq1_xmx_bytes(const ggml_tensor * tensor) { - return (size_t) (ggml_nelements(tensor) / QK_PTQ1_0) * sizeof(block_pq2_0); +// The XMX layout of PQ2_0/PTQ1_0 weights pads the rows to 16 (pq2_xmx.hpp), so on devices that run that path +// their allocation reserves the room +static bool ggml_sycl_xmx_pads(int device, const ggml_tensor * tensor) { + return (tensor->type == GGML_TYPE_PQ2_0 || tensor->type == GGML_TYPE_PTQ1_0) && g_ggml_sycl_enable_optimize && + !ggml_sycl_xmx_disabled() && tensor->ne[2] == 1 && tensor->ne[3] == 1 && + ggml_sycl_pq2_xmx_supports_ne0(tensor->ne[0]) && ggml_sycl_device_has_dpas16(device); } static size_t ggml_backend_sycl_buffer_type_get_alloc_size(ggml_backend_buffer_type_t buft, const ggml_tensor * tensor) { @@ -1006,8 +1002,8 @@ static size_t ggml_backend_sycl_buffer_type_get_alloc_size(ggml_backend_buffer_t } const auto * buft_ctx = (const ggml_backend_sycl_buffer_type_context *) buft->context; - if (ggml_sycl_ptq1_xmx_expands(buft_ctx->device, tensor)) { - size = std::max(size, ggml_sycl_ptq1_xmx_bytes(tensor)); + if (ggml_sycl_xmx_pads(buft_ctx->device, tensor)) { + size = std::max(size, ggml_sycl_pq2_xmx_bytes(tensor)); } return size; @@ -3808,7 +3804,7 @@ inline bool ggml_sycl_supports_mmq(enum ggml_type type) { return false; } -// The PQ2_0/PTQ1_0 XMX path feeds 2-bit weights to ESIMD DPAS at execution size 16 through 2D block loads, which +// The PQ2_0/PTQ1_0 XMX path feeds 2-bit weights to DPAS at execution size 16 through 2D block loads, which // every XMX device with 16-wide DPAS has (Xe-HPC, Xe2 and later). 8-wide XMX (Xe-HPG, Arrow Lake-H) keeps the // existing paths. static bool ggml_sycl_device_has_dpas16(int device) { @@ -4612,9 +4608,8 @@ static bool ggml_sycl_pq2_xmx_use(ggml_backend_sycl_context & ctx, const ggml_te !ggml_is_contiguous(dst)) { return false; } - // PTQ1_0 expands to 34 bytes a block, which only fits where the buffer reserved room for it - if (src0->type == GGML_TYPE_PTQ1_0 && - ggml_backend_buft_get_alloc_size(src0->buffer->buft, src0) < ggml_sycl_ptq1_xmx_bytes(src0)) { + // the padded rows only fit where the buffer reserved room for them + if (ggml_backend_buft_get_alloc_size(src0->buffer->buft, src0) < ggml_sycl_pq2_xmx_bytes(src0)) { return false; } if (!ggml_sycl_pq2_xmx_reorder(const_cast(src0), ctx.stream())) { diff --git a/ggml/src/ggml-sycl/pq2_xmx.cpp b/ggml/src/ggml-sycl/pq2_xmx.cpp index 20fbf60f4ae9..5b8453190939 100644 --- a/ggml/src/ggml-sycl/pq2_xmx.cpp +++ b/ggml/src/ggml-sycl/pq2_xmx.cpp @@ -1,368 +1,1107 @@ +// +// MIT license +// Copyright (C) 2024 Intel Corporation +// SPDX-License-Identifier: MIT +// +// The Xe2 helpers and the GEMV / GEMM kernels are ported from TernSYCL int2_via_int2_x_int8_dpas +// (https://github.com/libxsmm/TernSYCL), distributed under this license: +// +// BSD 3-Clause License +// +// Copyright (c) 2026, Intel Corporation +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions are met: +// +// * Redistributions of source code must retain the above copyright notice, this +// list of conditions and the following disclaimer. +// +// * Redistributions in binary form must reproduce the above copyright notice, +// this list of conditions and the following disclaimer in the documentation +// and/or other materials provided with the distribution. +// +// * Neither the name of the copyright holder nor the names of its +// contributors may be used to endorse or promote products derived from +// this software without specific prior written permission. +// +// THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +// AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +// IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +// DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +// FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +// SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +// CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +// OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +// OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +// + #include "pq2_xmx.hpp" -#include "dequantize.hpp" -// GGML_SYCL_NO_PQ2_XMX: an AOT build for a device without 16-wide DPAS, which cannot compile these kernels #if defined(__INTEL_LLVM_COMPILER) && !defined(GGML_SYCL_NO_PQ2_XMX) -#include +#include + +#include +#include + +// Kernel functor types must not be in an anonymous namespace (SYCL kernel names). +namespace ggml_sycl_xmx { + +namespace syclex = sycl::ext::oneapi::experimental; +namespace intelex = sycl::ext::intel::experimental; + +#define XE2_VEC(T, n, name) typedef T name __attribute__((ext_vector_type(n))) +XE2_VEC(short, 2, short2); +XE2_VEC(short, 4, short4); +XE2_VEC(short, 8, short8); +XE2_VEC(unsigned short, 2, ushort2); +XE2_VEC(unsigned short, 4, ushort4); +XE2_VEC(unsigned short, 16, ushort16); +XE2_VEC(int, 2, int2); +XE2_VEC(int, 4, int4); +XE2_VEC(int, 8, int8); +XE2_VEC(unsigned, 8, uint8); +XE2_VEC(float, 2, float2); +XE2_VEC(float, 4, float4); +XE2_VEC(float, 8, float8); +#undef XE2_VEC + +constexpr int GS = QK_PQ2_0; +constexpr float EPS = 1.1920928955078125e-07f; // FLT_EPSILON + +static_assert(GS == 128 && sizeof(block_pq2_0) == 34, "PQ2_0 layout changed"); + +// SA row pitch, padded so SA is a valid 2D block surface +inline int ldsa(int M) { return (M + 31) & ~31; } + +// 2D surface: width and pitch in bytes, height in rows +struct surf { + long long base; + int w, h, p; + surf(const void * b, int width_bytes, int height, int pitch_bytes) : + base((long long) b), w(width_bytes - 1), h(height - 1), p(pitch_bytes - 1) {} +}; + +#ifdef __SYCL_DEVICE_ONLY__ +#define XE2_ASM(...) __asm__(__VA_ARGS__) +#define XE2_ASM_V(...) __asm__ volatile(__VA_ARGS__) +#else +#define XE2_ASM(...) +#define XE2_ASM_V(...) +#endif + +// lsc 2D block read / write on an explicit surface: flat[base, width-1, height-1, pitch-1, x, y] +inline uint8 rd_32b_8r16(const surf & s, int x, int y) { + uint8 v; + XE2_ASM("{\n" + ".decl SB v_type=G type=q num_elts=1 align=qword alias=<%1,0>\n" + ".decl SW v_type=G type=d num_elts=1 align=dword alias=<%2,0>\n" + ".decl SH v_type=G type=d num_elts=1 align=dword alias=<%3,0>\n" + ".decl SP v_type=G type=d num_elts=1 align=dword alias=<%4,0>\n" + ".decl SX v_type=G type=d num_elts=1 align=dword alias=<%5,0>\n" + ".decl SY v_type=G type=d num_elts=1 align=dword alias=<%6,0>\n" + "lsc_load_block2d.ugm (M1, 1) %0:d32.16x8nn flat[SB,SW,SH,SP,SX,SY]\n" + "}\n" + : "=rw"(v) + : "rw.u"(s.base), "rw.u"(s.w), "rw.u"(s.h), "rw.u"(s.p), "rw.u"(x), "rw.u"(y)); + return v; +} + +inline void wr_32b_8r16(const surf & s, int x, int y, uint8 v) { + XE2_ASM_V("{\n" + ".decl SB v_type=G type=q num_elts=1 align=qword alias=<%0,0>\n" + ".decl SW v_type=G type=d num_elts=1 align=dword alias=<%1,0>\n" + ".decl SH v_type=G type=d num_elts=1 align=dword alias=<%2,0>\n" + ".decl SP v_type=G type=d num_elts=1 align=dword alias=<%3,0>\n" + ".decl SX v_type=G type=d num_elts=1 align=dword alias=<%4,0>\n" + ".decl SY v_type=G type=d num_elts=1 align=dword alias=<%5,0>\n" + "lsc_store_block2d.ugm (M1, 1) flat[SB,SW,SH,SP,SX,SY] %6:d32.16x8nn\n" + "}\n" + :: "rw.u"(s.base), "rw.u"(s.w), "rw.u"(s.h), "rw.u"(s.p), "rw.u"(x), "rw.u"(y), "rw"(v)); +} + +// 2D block reads from a prebuilt address payload plus immediate (DX, DY) offsets. +// The vISA text is built at compile time. +namespace detail { +template struct cstr { + char s[N]{}; + size_t n = 0; + constexpr size_t size() const { return n; } + constexpr const char * data() const { return s; } + constexpr void add(const char * p) { + while (*p) { + s[n++] = *p++; + } + } + constexpr void addi(int v) { + if (v < 0) { + s[n++] = '-'; + v = -v; + } + char t[12]{}; + int k = 0; + do { + t[k++] = char('0' + v % 10); + v /= 10; + } while (v); + while (k) { + s[n++] = t[--k]; + } + } +}; + +template constexpr auto rd2d_str() { + cstr<320> c; + c.add("{\n.decl PD v_type=G type=ud num_elts=8 align=GRF alias=<%1,0>\n"); + if constexpr (SH::pad) { + c.add(".decl TP v_type=G type=uw num_elts=32 align=GRF\nlsc_load_block2d.ugm (M1, 1) TP:"); + } else { + c.add("lsc_load_block2d.ugm (M1, 1) %0:"); + } + c.add(SH::v); + c.add(" flat[PD + ("); + c.addi(DX); + c.add(","); + c.addi(DY); + c.add(")]\n"); + if constexpr (SH::pad) { + c.add("mov (M1, 16) %0(0,0)<1> TP(0,0)<1;1,0>\n"); + } + c.add("}\n"); + return c; +} +} // namespace detail + +// block shape: vISA type and payload dword 7 = (V-1) << 16 | (R-1) << 8 | (C-1). +// pad: the block is half a GRF but the load writes a whole GRF. +struct b32_16x1 { static constexpr const char * v = "d32.16x1nn"; static constexpr int code = 0x00f; static constexpr bool pad = false; }; +struct b32_16x8 { static constexpr const char * v = "d32.16x8nn"; static constexpr int code = 0x70f; static constexpr bool pad = false; }; +struct b16_2x16x8 { static constexpr const char * v = "d16.2x16x8nn"; static constexpr int code = 0x1070f; static constexpr bool pad = false; }; +struct b16_32x1 { static constexpr const char * v = "d16.32x1nn"; static constexpr int code = 0x01f; static constexpr bool pad = false; }; +struct b16_16x1 { static constexpr const char * v = "d16.16x1nn"; static constexpr int code = 0x00f; static constexpr bool pad = true; }; + +template inline unsigned pl2d(const surf & s, int x, int y) { + unsigned pl; + XE2_ASM("{\n" + ".decl PQ v_type=G type=uq num_elts=4 align=GRF alias=<%0,0>\n" + ".decl PD v_type=G type=ud num_elts=8 align=GRF alias=<%0,0>\n" + "mov (M1_NM, 1) PQ(0,0)<1> %1(0,0)<0;1,0>\n" + "mov (M1_NM, 1) PD(0,2)<1> %2(0,0)<0;1,0>\n" + "mov (M1_NM, 1) PD(0,3)<1> %3(0,0)<0;1,0>\n" + "mov (M1_NM, 1) PD(0,4)<1> %4(0,0)<0;1,0>\n" + "mov (M1_NM, 1) PD(0,5)<1> %5(0,0)<0;1,0>\n" + "mov (M1_NM, 1) PD(0,6)<1> %6(0,0)<0;1,0>\n" + "mov (M1_NM, 1) PD(0,7)<1> %7(0,0)<0;1,0>\n" + "}\n" + : "=rw"(pl) + : "rw.u"(s.base), "rw.u"(s.w), "rw.u"(s.h), "rw.u"(s.p), "rw.u"(x), "rw.u"(y), "rw.u"(SH::code)); + return pl; +} + +// move the block x (dword 5) or y (dword 6) of a payload in place +inline void pl2d_x(unsigned & pl, int v) { + XE2_ASM("{\n.decl PD v_type=G type=ud num_elts=8 align=GRF alias=<%0,0>\n" + "mov (M1_NM, 1) PD(0,5)<1> %1(0,0)<0;1,0>\n}\n" : "+rw"(pl) : "rw.u"(v)); +} + +inline void pl2d_y(unsigned & pl, int v) { + XE2_ASM("{\n.decl PD v_type=G type=ud num_elts=8 align=GRF alias=<%0,0>\n" + "mov (M1_NM, 1) PD(0,6)<1> %1(0,0)<0;1,0>\n}\n" : "+rw"(pl) : "rw.u"(v)); +} + +template inline T rd2d(unsigned pl) { + T v; + XE2_ASM((detail::rd2d_str()) : "=rw"(v) : "rw"(pl)); + return v; +} -namespace { +template inline void static_for_impl(F && f, std::integer_sequence) { + (f(std::integral_constant{}), ...); +} -namespace esimd = sycl::ext::intel::esimd; -namespace xmx = sycl::ext::intel::esimd::xmx; +template inline void static_for(F && f) { + static_for_impl(f, std::make_integer_sequence{}); +} -constexpr int PQ2_XMX_QS_BYTES = QK_PQ2_0 / 4; // 32 bytes of 2-bit codes per block -constexpr int PQ2_XMX_QS_DW = PQ2_XMX_QS_BYTES / 4; -constexpr int PQ2_XMX_WG = 16; // independent tiles per work-group when K is not split +// Sub-group block reads, striped: element i of lane l = p[l + 16 i]. p must be uniform and 4-byte aligned. +inline unsigned short sg_rd_us(const unsigned short * p) { + unsigned short v; + XE2_ASM("{\n.decl TP v_type=G type=uw num_elts=32 align=GRF\n" + ".decl AD v_type=G type=q num_elts=1 align=GRF\n" + "mov (M1_NM, 1) AD(0,0)<1> %1(0,0)<0;1,0>\n" + "lsc_load.ugm (M1_NM, 1) TP:d32x8t flat[AD]:a64\n" + "mov (M1, 16) %0(0,0)<1> TP(0,0)<1;1,0>\n}\n" + : "=rw"(v) : "rw.u"((long long) p)); + return v; +} -// thread count the K split aims for: the B50 runs 1024 hardware threads, and decode needs several in flight -// per EU to keep enough loads outstanding -constexpr int PQ2_XMX_TARGET_THREADS = 4096; -constexpr int PQ2_XMX_PREFILL_TARGET_THREADS = 512; +inline ushort4 sg_rd_us4(const unsigned short * p) { + ushort4 v; + XE2_ASM("{\n.decl AD v_type=G type=q num_elts=1 align=GRF\n" + "mov (M1_NM, 1) AD(0,0)<1> %1(0,0)<0;1,0>\n" + "lsc_load.ugm (M1_NM, 1) %0:d32x32t flat[AD]:a64\n}\n" + : "=rw"(v) : "rw.u"((long long) p)); + return v; +} -static_assert(QK_PQ2_0 == 128 && sizeof(block_pq2_0) == 34, "PQ2_0 layout changed"); -static_assert(QK_PTQ1_0 == QK_PQ2_0, "PTQ1_0 expands block for block into PQ2_0 codes"); +// DPAS s8 A (K32, one short per row per lane) x s2 B (two dwords per lane), systolic depth 8, R rows of A +#define XE2_DPAS(A, C, R, na) \ + inline C dpas_s2s8(A a, int2 b, C acc) { \ + XE2_ASM("{\n" \ + ".decl DB v_type=G type=ud num_elts=32 align=GRF alias=<%1,0>\n" \ + ".decl DA v_type=G type=ud num_elts=" #na " align=GRF alias=<%2,0>\n" \ + "dpas.s2.s8.8." #R " (M1, 16) %0.0 %0.0 DB.0 DA(0,0)\n" \ + "}\n" \ + : "+rw"(acc) : "rw"(b), "rw"(a)); \ + return acc; \ + } \ + inline C dpas_s2s8_z(A a, int2 b) { \ + C d; \ + XE2_ASM("{\n" \ + ".decl DB v_type=G type=ud num_elts=32 align=GRF alias=<%1,0>\n" \ + ".decl DA v_type=G type=ud num_elts=" #na " align=GRF alias=<%2,0>\n" \ + "dpas.s2.s8.8." #R " (M1, 16) %0.0 %%null.0 DB.0 DA(0,0)\n" \ + "}\n" \ + : "=rw"(d) : "rw"(b), "rw"(a)); \ + return d; \ + } +XE2_DPAS(short, int, 1, 8) +XE2_DPAS(short2, int2, 2, 16) +XE2_DPAS(short4, int4, 4, 32) +XE2_DPAS(short8, int8, 8, 64) +#undef XE2_DPAS + +inline float half_bits_to_float(unsigned short h) { + return (float) sycl::bit_cast(h); +} -// PQ2_0 packs value+1 per 2-bit field, lowest first: as a little-endian dword that is already DPAS 2-bit packing. -// Subtracting 1 per field without borrow gives the s2 values DPAS multiplies. -template -ESIMD_INLINE esimd::simd pq2_codes_to_s2(esimd::simd x) { +// PQ2_0 stores value+1 in each 2-bit field. Subtract 1 per field without borrow across fields +// to get s2 two's complement. +inline uint32_t pq2_codes_to_s2(uint32_t x) { constexpr uint32_t H = 0xAAAAAAAAu; constexpr uint32_t L = 0x55555555u; return ((x | H) - L) ^ (~x & H); } -// Each thread computes an (8*MR) token x (16*NR) row tile over its share of the K blocks. With S > 1 the S -// threads of a work-group split K for the same tile and reduce through SLM, so a mat-vec keeps enough threads -// streaming weights. One DPAS covers 8 tokens x 16 rows x 32 k; the four of a block accumulate in int32. -template -ESIMD_INLINE void pq2_xmx_thread(const uint32_t * wq, const uint16_t * wd, const uint32_t * a8, const float * as, - float * dst, int K, int nrows, int ncols, int nrows_dst, int n_tiles_n, - int tile, int ks) { - using namespace esimd; - constexpr int TM = 8 * MR; - constexpr int TN = 16 * NR; +// weight formats of the reordered layouts +enum { FMT_PQ2 = 0, FMT_PTQ1 = 1 }; - if constexpr (S > 1) { - slm_init(); +// PTQ1_0 packs 5 trits per byte. The kernels decode them in byte-major K order (the trits of one +// byte next to each other): then the 10-bit LUT entries of consecutive bytes concatenate into the +// s2 dwords. The activation is quantized in the same order, see ptq1_k_orig(). +constexpr int PTQ1_ROWS = 7; // dwords per 128-group and column after the reorder: qs[24], qh[2] + d + +inline int ptq1_trit(int b, int n) { + constexpr uint8_t pow3[5] = { 1, 3, 9, 27, 81 }; + const uint8_t q = (uint8_t) (b * pow3[n]); + return ((int) q * 3) >> 8; +} + +// the s2 codes of the 5 trits of byte b, trit n at bits 2n +inline uint16_t ptq1_lut_entry(int b) { + uint16_t e = 0; + for (int n = 0; n < 5; ++n) { + e |= ((ptq1_trit(b, n) + 3) & 3) << (2 * n); } + return e; +} - const int m0 = (tile / n_tiles_n) * TM; - const int n0 = (tile % n_tiles_n) * TN; - const int nb = K / QK_PQ2_0; - const int b0 = (int) ((int64_t) ks * nb / S); - const int b1 = (int) ((int64_t) (ks + 1) * nb / S); - - // weights: nrows rows of nb*32 bytes; activations: ncols rows of K bytes. Rows past either read as zeros. - const uint32_t wsurf_w = (uint32_t) (nb * PQ2_XMX_QS_BYTES) - 1; - const uint32_t wsurf_h = (uint32_t) nrows - 1; - const uint32_t asurf_w = (uint32_t) K - 1; - const uint32_t asurf_h = (uint32_t) ncols - 1; - - // scale gathers clamp to a valid row/token; the clamped lanes are never stored - simd d_off[NR]; -#pragma unroll - for (int g = 0; g < NR; ++g) { - simd r(n0 + 16 * g, 1); - r.merge(simd(nrows - 1), r >= (uint32_t) nrows); - d_off[g] = r * (uint32_t) (nb * sizeof(uint16_t)); +// position p of the decode order -> index of that weight in the PTQ1_0 block +inline int ptq1_k_orig(int p) { + if (p < 80) { + return (p % 5) * 16 + p / 5; } - simd s_off[MR]; -#pragma unroll - for (int s = 0; s < MR; ++s) { - simd m(m0 + 8 * s, 1); - m.merge(simd(ncols - 1), m >= (uint32_t) ncols); - s_off[s] = m * (uint32_t) (nb * sizeof(float)); + if (p < 120) { + return 80 + ((p - 80) % 5) * 8 + (p - 80) / 5; } + return 120 + ((p - 120) % 4) * 2 + (p - 120) / 4; +} - simd acc[MR][NR]; +// 7 raw dwords of a 128-group (r[7] unused) -> 8 s2 dwords, 16 weights each +inline uint8 ptq1_decode(const uint8 & r, const uint16_t * lut) { + uint32_t e[16]; + uint32_t f[8]; #pragma unroll - for (int s = 0; s < MR; ++s) { + for (int m = 0; m < 16; ++m) { + e[m] = lut[(r[m / 4] >> (8 * (m % 4))) & 0xFF]; + } #pragma unroll - for (int g = 0; g < NR; ++g) { - acc[s][g] = 0.0f; - } + for (int m = 0; m < 8; ++m) { + f[m] = lut[(r[4 + m / 4] >> (8 * (m % 4))) & 0xFF]; } + const uint32_t g0 = lut[r[6] & 0xFF] & 0xFF; + const uint32_t g1 = lut[(r[6] >> 8) & 0xFF] & 0xFF; + + uint8 w; + w[0] = e[0] | e[1] << 10 | e[2] << 20 | e[3] << 30; + w[1] = e[3] >> 2 | e[4] << 8 | e[5] << 18 | e[6] << 28; + w[2] = e[6] >> 4 | e[7] << 6 | e[8] << 16 | e[9] << 26; + w[3] = e[9] >> 6 | e[10] << 4 | e[11] << 14 | e[12] << 24; + w[4] = e[12] >> 8 | e[13] << 2 | e[14] << 12 | e[15] << 22; + w[5] = f[0] | f[1] << 10 | f[2] << 20 | f[3] << 30; + w[6] = f[3] >> 2 | f[4] << 8 | f[5] << 18 | f[6] << 28; + w[7] = f[6] >> 4 | f[7] << 6 | g0 << 16 | g1 << 24; + return w; +} + +// fill the decode table in local memory; all work-items of the group must call it +inline uint16_t * ptq1_lut(sycl::nd_item<2> it) { + uint16_t * lut = *sycl::ext::oneapi::group_local_memory_for_overwrite(it.get_group()); + for (int i = it.get_local_linear_id(); i < 256; i += it.get_local_range().size()) { + lut[i] = ptq1_lut_entry(i); + } + sycl::group_barrier(it.get_group()); + return lut; +} + +// One sub-group per (row, 128-group): SA = 127 / absmax, Aq = rint(A * SA). Rows of src1 are +// flattened over dims 1..3. +struct quant_a { + const float * A; + float * SA; + int8_t * Aq; + int M, K, ne11, ne12; + int64_t s11, s12, s13; + int perm; // write each 128-group in PTQ1_0 decode order - for (int b = b0; b < b1; ++b) { - // transposed load: w[g][j*16 + n] = dword j (k = 16j..16j+15) of row n, the DPAS B layout for 2 dwords per k32 - simd w[NR]; - simd dw[NR]; -#pragma unroll - for (int g = 0; g < NR; ++g) { - w[g] = pq2_codes_to_s2<128>(load_2d( - wq, wsurf_w, wsurf_h, wsurf_w, b * PQ2_XMX_QS_DW, n0 + 16 * g)); - simd dbits = gather(wd, d_off[g] + (uint32_t) (b * sizeof(uint16_t))); - simd dh = dbits.template bit_cast_view(); - dw[g] = convert(dh); + void operator()(sycl::nd_item<2> it) const { + const auto sg = it.get_sub_group(); + const int lane = sg.get_local_linear_id(); + const int g = (int) it.get_global_id(1) / 16; + const int m = (int) it.get_global_id(0); + if (g >= K / GS || m >= M) { + return; } + const int i1 = m % ne11; + const int i2 = (m / ne11) % ne12; + const int i3 = m / (ne11 * ne12); + const float * a = A + i1 * s11 + i2 * s12 + i3 * s13 + g * GS + 8 * lane; + float v[8]; + float mx = 0.0f; +#pragma unroll + for (int i = 0; i < 8; ++i) { + v[i] = a[i]; + mx = sycl::fmax(mx, sycl::fabs(v[i])); + } + mx = sycl::reduce_over_group(sg, mx, sycl::maximum()); + const float s = 127.0f / sycl::fmax(mx, EPS); + if (lane == 0) { + SA[(size_t) g * ldsa(M) + m] = s; + } + if (perm) { +#pragma unroll + for (int i = 0; i < 8; ++i) { + v[i] = a[ptq1_k_orig(8 * lane + i) - 8 * lane]; + } + } + uint64_t q = 0; #pragma unroll - for (int s = 0; s < MR; ++s) { - const int yrow = m0 + 8 * s; + for (int i = 0; i < 8; ++i) { + q |= (uint64_t) (uint8_t) (int8_t) sycl::clamp(sycl::rint(v[i] * s), -128.0f, 127.0f) << (8 * i); + } + *(uint64_t *) (Aq + (size_t) m * K + g * GS + 8 * lane) = q; + } + + auto get(syclex::properties_tag) const { + return syclex::properties{ syclex::sub_group_size<16>, syclex::work_group_size<1, 16> }; + } +}; + +template struct rows; +template <> struct rows<1> { using a_t = short; using ia_t = int; using fa_t = float; }; +template <> struct rows<2> { using a_t = short2; using ia_t = int2; using fa_t = float2; }; +template <> struct rows<4> { using a_t = short4; using ia_t = int4; using fa_t = float4; }; +template <> struct rows<8> { using a_t = short8; using ia_t = int8; using fa_t = float8; }; + +template inline auto el(const V & v, int r) { + if constexpr (SGM == 1) { + return v; + } else { + return v[r]; + } +} + +template inline void set_el(V & v, int r, T x) { + if constexpr (SGM == 1) { + v = x; + } else { + v[r] = x; + } +} + +// GEMV / small M. A sub-group owns 16 columns and SGM rows and walks its K slice in 128-steps: +// one 2D read of 8 B dwords (4 DPAS of K = 32), one SB row, SGM A rows. +// NSG_N sub-groups along N, LS K-slices per column block (reduced in SLM), +// U 128-steps whose loads are issued before any compute. +template struct gemv { + const int8_t * Aq; + const float * SA; + const uint32_t * B; + const unsigned short * SB; + float * C; + int M, N, NP, K, ldc; - simd ci[NR]; + using a_t = typename rows::a_t; + using ia_t = typename rows::ia_t; + using fa_t = typename rows::fa_t; + static constexpr int WG = 16 * NSG_N * LS; + + // SGM x 128 A tile of step s: aq[c] = K 32c..32c+31, inv[r] = 1 / SA of row r + void load_a(int m0, int s, a_t * aq, float * inv) const { + const int lda = ldsa(M); +#pragma unroll + for (int r = 0; r < SGM; ++r) { + const bool ok = m0 + r < M; + const size_t row = (size_t) sycl::min(m0 + r, M - 1) * K + s * GS; + const ushort4 l = sg_rd_us4((const unsigned short *) (Aq + row)); + const ushort4 v = ok ? l : ushort4{}; #pragma unroll - for (int g = 0; g < NR; ++g) { - ci[g] = 0; + for (int c = 0; c < 4; ++c) { + set_el(aq[c], r, (short) v[c]); } + inv[r] = ok ? sycl::native::recip(SA[(size_t) s * lda + m0 + r]) : 0.0f; + } + } + + static fa_t step(fa_t acc, const uint8 & w, float sb, const a_t * aq, const float * inv) { + ia_t ia = dpas_s2s8_z(aq[0], int2{ (int) w[0], (int) w[1] }); #pragma unroll - for (int c = 0; c < QK_PQ2_0 / 32; ++c) { - // A operand: token t's 32 int8 values at dwords t*8 .. t*8+7 - simd ad = load_2d(a8, asurf_w, asurf_h, asurf_w, - b * 32 + 8 * c, yrow); - simd am = ad.template bit_cast_view(); -#pragma unroll - for (int g = 0; g < NR; ++g) { - simd bd = w[g].template select<32, 1>(32 * c); - simd bm = bd.template bit_cast_view(); - ci[g] = xmx::dpas<8, 8, int, int, signed char, signed char, xmx::dpas_argument_type::s2, - xmx::dpas_argument_type::s8>(ci[g], bm, am); - } + for (int c = 1; c < 4; ++c) { + ia = dpas_s2s8(aq[c], int2{ (int) w[2 * c], (int) w[2 * c + 1] }, ia); + } +#pragma unroll + for (int r = 0; r < SGM; ++r) { + set_el(acc, r, el(acc, r) + (float) el(ia, r) * (sb * inv[r])); + } + return acc; + } + + void operator()(sycl::nd_item<2> it) const { + const auto sgp = it.get_sub_group(); + const int lane = sgp.get_local_linear_id(); + const int sg = sgp.get_group_linear_id(); + const int sgn = sg % NSG_N; + const int sgk = sg / NSG_N; + const int n0 = ((int) it.get_group(1) * NSG_N + sgn) * 16; + const int m0 = (int) it.get_group(0) * SGM; + + const int nsteps = K / GS; + const int per = (nsteps + LS - 1) / LS; + const int s_begin = sgk * per; + const int s_end = sycl::min(nsteps, s_begin + per); + const int rows = FMT == FMT_PTQ1 ? PTQ1_ROWS : 8; // B rows per 128-group + const surf sbs(B, NP * 4, nsteps * rows, NP * 4); + + const uint16_t * lut = nullptr; + if constexpr (FMT == FMT_PTQ1) { + lut = ptq1_lut(it); + } + // B of step s: s2 codes, and the scale (stored with the codes for PTQ1_0) + auto load_w = [&](int s, uint8 & w, float & sb) { + w = rd_32b_8r16(sbs, n0, s * rows); + if constexpr (FMT == FMT_PQ2) { + sb = half_bits_to_float(sg_rd_us(SB + (size_t) s * NP + n0)); } + }; + auto decode = [&](uint8 & w, float & sb) { + if constexpr (FMT == FMT_PTQ1) { + sb = half_bits_to_float((unsigned short) (w[6] >> 16)); + w = ptq1_decode(w, lut); + } + }; - const simd da = gather(as, s_off[s] + (uint32_t) (b * sizeof(float))); + fa_t acc = 0.0f; + if (n0 < N) { + int s = s_begin; +#pragma unroll 1 + for (; s + U <= s_end; s += U) { + uint8 w[U]; + float sb[U]; + a_t aq[U][4]; + float inv[U][SGM]; #pragma unroll - for (int g = 0; g < NR; ++g) { + for (int u = 0; u < U; ++u) { + load_w(s + u, w[u], sb[u]); + load_a(m0, s + u, aq[u], inv[u]); + } #pragma unroll - for (int t = 0; t < 8; ++t) { - const simd cit = ci[g].template select<16, 1>(16 * t); - acc[s][g].template select<16, 1>(16 * t) += convert(cit) * (dw[g] * da[t]); + for (int u = 0; u < U; ++u) { + decode(w[u], sb[u]); + acc = step(acc, w[u], sb[u], aq[u], inv[u]); + } + } + if constexpr (U > 1) { + for (; s < s_end; ++s) { + a_t aq[4]; + float inv[SGM]; + uint8 w; + float sb; + load_w(s, w, sb); + load_a(m0, s, aq, inv); + decode(w, sb); + acc = step(acc, w, sb, aq, inv); } } } - } - if constexpr (S > 1) { - // every thread parks its partial tile in SLM; thread 0 sums them and stores -#pragma unroll - for (int s = 0; s < MR; ++s) { + if constexpr (LS > 1) { + float * red = *sycl::ext::oneapi::group_local_memory_for_overwrite( + it.get_group()); + if (sgk > 0) { + float * dst = red + (((sgk - 1) * NSG_N + sgn) * SGM) * 16; #pragma unroll - for (int g = 0; g < NR; ++g) { + for (int r = 0; r < SGM; ++r) { + dst[r * 16 + lane] = el(acc, r); + } + } + sycl::group_barrier(it.get_group()); + if (sgk > 0) { + return; + } + for (int j = 0; j < LS - 1; ++j) { + const float * src = red + ((j * NSG_N + sgn) * SGM) * 16; #pragma unroll - for (int q = 0; q < 8; ++q) { - const uint32_t off = (uint32_t) ((((ks * MR + s) * NR + g) * 128 + 16 * q) * sizeof(float)); - slm_block_store(off, acc[s][g].template select<16, 1>(16 * q)); + for (int r = 0; r < SGM; ++r) { + set_el(acc, r, el(acc, r) + src[r * 16 + lane]); } } } - barrier(); - if (ks != 0) { + + if (n0 + lane >= N) { return; } - for (int o = 1; o < S; ++o) { #pragma unroll - for (int s = 0; s < MR; ++s) { + for (int r = 0; r < SGM; ++r) { + if (m0 + r < M) { + C[(size_t) (m0 + r) * ldc + n0 + lane] = el(acc, r); + } + } + } + + auto get(syclex::properties_tag) const { + return syclex::properties{ syclex::sub_group_size<16>, syclex::work_group_size<1, WG> }; + } +}; + +// Large-M GEMM. A sub-group computes an MT_M x MT_N tile, a work-group is WG_M x WG_N sub-groups. +// Per 128-group, B and SB are loaded once for all MT_M/8 row blocks and each 8 x 128 A block once +// for all MT_N/16 column blocks. 2D block I/O zero-fills out-of-range reads and clips writes. +template struct gemm { + const int8_t * Aq; + const float * SA; + const uint32_t * B; + const unsigned short * SB; + float * C; + int M, N, NP, K, ldc; + bool st2d; // C is a valid 2D surface + float * part; // with ks > 1: per K-slice results [ks][M][NP], summed by launch_gemm + int ks; + + static constexpr int MB = MT_M / 8, NB = MT_N / 16, WG = 16 * WG_M * WG_N; + + void operator()(sycl::nd_item<2> it) const { + const auto sgp = it.get_sub_group(); + const int sg = sgp.get_group_linear_id(); + const int kz = (int) it.get_group(0) % ks; // K slice of this work-group + const int m0 = ((int) it.get_group(0) / ks * WG_M + sg / WG_N) * MT_M; + const int n0 = ((int) it.get_group(1) * WG_N + sg % WG_N) * MT_N; + const int rows = FMT == FMT_PTQ1 ? PTQ1_ROWS : 8; // B rows per 128-group + const surf sbs(B, NP * 4, K / GS * rows, NP * 4); + + const uint16_t * lut = nullptr; + if constexpr (FMT == FMT_PTQ1) { + lut = ptq1_lut(it); + } + const surf ssb(SB, NP * 2, K / GS, NP * 2); + // lane l gets SA[s, m + l]; pad columns past M are junk + const surf ssa(SA, ldsa(M) * 4, K / GS, ldsa(M) * 4); + + float8 acc[MB][NB]; +#pragma unroll + for (int i = 0; i < MB; ++i) { +#pragma unroll + for (int j = 0; j < NB; ++j) { + acc[i][j] = 0.0f; + } + } + + // 2D payloads built once; each K step only moves y (B, SB, SA) or x (A) + unsigned pb = pl2d(sbs, n0, 0); + unsigned psb2 = pl2d(ssb, n0, 0); + unsigned psb1 = pl2d(ssb, n0, 0); + unsigned psa = pl2d(ssa, m0, 0); + unsigned pq = pl2d(surf(Aq, K, M, K), 0, m0); + + // one 128-step: A for all row blocks, B and its scales through get_b + auto step = [&](int s, auto && get_b) { + // A and SA loads first, they feed the first dpas. Row block 0 before B and SB, + // block I + 1 at the start of block I. + pl2d_y(psa, s); + pl2d_x(pq, s * GS / 2); + unsigned sar[MB]; + ushort16 ar[MB][2]; + auto load_a = [&](auto ii) { + constexpr int I = decltype(ii)::value; + sar[I] = rd2d(psa); + static_for<2>([&](auto h) { + ar[I][h] = rd2d(pq); + }); + }; + load_a(std::integral_constant{}); + uint8 w[NB]; + float sb[NB]; + get_b(s, w, sb); + static_for([&](auto ii) { + constexpr int I = decltype(ii)::value; + if constexpr (I + 1 < MB) { + load_a(std::integral_constant{}); + } + short8 aq[4]; + float inv[8]; +#pragma unroll + for (int h = 0; h < 2; ++h) { +#pragma unroll + for (int r = 0; r < 8; ++r) { + aq[2 * h][r] = (short) ar[I][h][r]; + aq[2 * h + 1][r] = (short) ar[I][h][8 + r]; + } + } + // rows >= M get junk here, but their int32 dot is 0 and the store clips them + const float invl = sycl::native::recip(sycl::bit_cast(sar[I])); +#pragma unroll + for (int r = 0; r < 8; ++r) { + inv[r] = sycl::group_broadcast(sgp, invl, r); + } +#pragma unroll + for (int j = 0; j < NB; ++j) { + int8 ia = dpas_s2s8_z(aq[0], int2{ (int) w[j][0], (int) w[j][1] }); #pragma unroll - for (int g = 0; g < NR; ++g) { + for (int c = 1; c < 4; ++c) { + ia = dpas_s2s8(aq[c], int2{ (int) w[j][2 * c], (int) w[j][2 * c + 1] }, ia); + } + // whole-vector convert: per-element casts go through a scratch register + const float8 fi = __builtin_convertvector(ia, float8); +#pragma unroll + for (int r = 0; r < 8; ++r) { + acc[I][j][r] += fi[r] * (sb[j] * inv[r]); + } + } + }); + }; + + const int nsteps = K / GS; + const int per = (nsteps + ks - 1) / ks; + const int kb = kz * per; + const int ke = sycl::min(nsteps, kb + per); + if constexpr (FMT == FMT_PQ2) { + for (int s = kb; s < ke; ++s) { + step(s, [&](int s, uint8 * w, float * sb) { + pl2d_y(pb, s * rows); + static_for([&](auto j) { w[j] = rd2d(pb); }); + // all SB loads first, then convert, so the loads do not serialize on one register + ushort2 sbr[(NB + 1) / 2]; + pl2d_y(psb2, s); + pl2d_y(psb1, s); + static_for<(NB + 1) / 2>([&](auto h) { + constexpr int J = 2 * decltype(h)::value; + if constexpr (J + 1 < NB) { + sbr[h] = rd2d(psb2); + } else { + sbr[h] = ushort2{ rd2d(psb1), 0 }; + } + }); +#pragma unroll + for (int j = 0; j < NB; ++j) { + sb[j] = half_bits_to_float(sbr[j / 2][j % 2]); + } + }); + } + } else if constexpr (WG_M == 1) { + // nothing to share: decode in registers + for (int s = kb; s < ke; ++s) { + step(s, [&](int s, uint8 * w, float * sb) { + pl2d_y(pb, s * rows); + static_for([&](auto j) { w[j] = rd2d(pb); }); #pragma unroll - for (int q = 0; q < 8; ++q) { - const uint32_t off = (uint32_t) ((((o * MR + s) * NR + g) * 128 + 16 * q) * sizeof(float)); - acc[s][g].template select<16, 1>(16 * q) += slm_block_load(off); + for (int j = 0; j < NB; ++j) { + sb[j] = half_bits_to_float((unsigned short) (w[j][6] >> 16)); + w[j] = ptq1_decode(w[j], lut); } + }); + } + } else { + // The WG_M sub-groups of a work-group column share their B columns: each decodes one of + // every WG_M steps into local memory, then all of them use the WG_M decoded steps. + const int lane = sgp.get_local_linear_id(); + const int wr = sg / WG_N; + const int wc = sg % WG_N; + uint32_t * dec = *sycl::ext::oneapi::group_local_memory_for_overwrite( + it.get_group()); + float * dsb = *sycl::ext::oneapi::group_local_memory_for_overwrite( + it.get_group()); + auto slot = [&](int u, int j) { return (u * WG_N + wc) * NB + j; }; + + for (int s0 = kb; s0 < ke; s0 += WG_M) { + if (s0 + wr < ke) { + pl2d_y(pb, (s0 + wr) * rows); + static_for([&](auto jj) { + constexpr int J = decltype(jj)::value; + const uint8 raw = rd2d(pb); + const uint8 w = ptq1_decode(raw, lut); +#pragma unroll + for (int d = 0; d < 8; ++d) { + dec[slot(wr, J) * 128 + d * 16 + lane] = w[d]; + } + dsb[slot(wr, J) * 16 + lane] = half_bits_to_float((unsigned short) (raw[6] >> 16)); + }); + } + sycl::group_barrier(it.get_group()); + const int nu = sycl::min(WG_M, ke - s0); + for (int u = 0; u < nu; ++u) { + step(s0 + u, [&](int, uint8 * w, float * sb) { +#pragma unroll + for (int j = 0; j < NB; ++j) { +#pragma unroll + for (int d = 0; d < 8; ++d) { + w[j][d] = dec[slot(u, j) * 128 + d * 16 + lane]; + } + sb[j] = dsb[slot(u, j) * 16 + lane]; + } + }); } + sycl::group_barrier(it.get_group()); } } - } - const simd lane(0, 1); + const int lane = sgp.get_local_linear_id(); + // a K slice writes its partial result; launch_gemm adds the slices up + float * out = ks > 1 ? part + (size_t) kz * M * NP : C; + const int ldo = ks > 1 ? NP : ldc; + const bool s2d = ks > 1 || st2d; + if (s2d) { + const surf sc(out, N * 4, M, ldo * 4); #pragma unroll - for (int s = 0; s < MR; ++s) { + for (int i = 0; i < MB; ++i) { #pragma unroll - for (int t = 0; t < 8; ++t) { - const int m = m0 + 8 * s + t; - if (m >= ncols) { - continue; + for (int j = 0; j < NB; ++j) { + wr_32b_8r16(sc, n0 + 16 * j, m0 + 8 * i, __builtin_bit_cast(uint8, acc[i][j])); + } } - // row base in 64 bits: an output head at a large ubatch passes 4 GB - float * drow = dst + (size_t) m * nrows_dst; + } else { +#pragma unroll + for (int i = 0; i < MB; ++i) { +#pragma unroll + for (int j = 0; j < NB; ++j) { + const int n = n0 + 16 * j + lane; #pragma unroll - for (int g = 0; g < NR; ++g) { - const simd n = lane + (uint32_t) (n0 + 16 * g); - const simd_mask<16> ok = n < (uint32_t) nrows; - scatter(drow, n * (uint32_t) sizeof(float), acc[s][g].template select<16, 1>(16 * t), ok); + for (int r = 0; r < 8; ++r) { + const int m = m0 + 8 * i + r; + if (m < M && n < N) { + out[(size_t) m * ldo + n] = acc[i][j][r]; + } + } + } } } } + + auto get(syclex::properties_tag) const { + return syclex::properties{ syclex::sub_group_size<16>, syclex::work_group_size<1, WG>, + intelex::grf_size<256> }; + } +}; + +struct args { + const int8_t * Aq; + const float * SA; + const uint32_t * B; + const unsigned short * SB; + float * C; + int M, N, NP, K, ldc; + bool st2d; + ggml_sycl_pool * pool; // for the K-slice partials + int threads; // hardware threads of the device +}; + +template static void launch_gemv(const args & a, dpct::queue_ptr stream) { + using kern = gemv; + const size_t wgn = 16 * NSG; + const sycl::range<2> local(1, kern::WG); + const sycl::range<2> global((a.M + SGM - 1) / SGM, (a.N + wgn - 1) / wgn * kern::WG); + stream->parallel_for(sycl::nd_range<2>(global, local), kern{ a.Aq, a.SA, a.B, a.SB, a.C, a.M, a.N, a.NP, a.K, a.ldc }); +} + +template static void launch_gemm(const args & a, dpct::queue_ptr stream) { + using kern = gemm; + const int mtiles = (a.M + MT_M * WG_M - 1) / (MT_M * WG_M); + const int ntiles = (a.N + MT_N * WG_N - 1) / (MT_N * WG_N); + // small batches leave most of the GPU idle: split K over work-groups and add the slices up after + const int ks = std::max(1, std::min(a.K / GS / 4, a.threads / (mtiles * ntiles * WG_M * WG_N))); + ggml_sycl_pool_alloc part(*a.pool); + if (ks > 1) { + part.alloc((size_t) ks * a.M * a.NP); + } + const sycl::range<2> local(1, kern::WG); + const sycl::range<2> global((size_t) mtiles * ks, (size_t) ntiles * kern::WG); + stream->parallel_for(sycl::nd_range<2>(global, local), kern{ a.Aq, a.SA, a.B, a.SB, a.C, a.M, a.N, a.NP, a.K, a.ldc, a.st2d, part.ptr, ks }); + if (ks > 1) { + const float * p = part.ptr; + float * C = a.C; + const int M = a.M, N = a.N, NP = a.NP, ldc = a.ldc; + stream->parallel_for(sycl::range<1>((size_t) M * N), [=](sycl::item<1> it) { + const int m = it[0] / N; + const int n = it[0] % N; + float v = 0.0f; + for (int k = 0; k < ks; ++k) { + v += p[((size_t) k * M + m) * NP + n]; + } + C[(size_t) m * ldc + n] = v; + }); + } } -template -static void launch_pq2_xmx(const uint32_t * wq, const uint16_t * wd, const uint32_t * a8, const float * as, - float * dst, int K, int nrows, int ncols, int nrows_dst, dpct::queue_ptr stream) { - constexpr int TM = 8 * MR; - constexpr int TN = 16 * NR; +template static void launch_gemv_ls(const args & a, int ls, dpct::queue_ptr stream) { + // PTQ1_0 decodes between the load and the dpas: more sub-groups per work-group hide that latency + constexpr int NSG = FMT == FMT_PTQ1 ? 4 : 2; + switch (ls) { + case 1: launch_gemv(a, stream); break; + case 2: launch_gemv(a, stream); break; + case 4: launch_gemv(a, stream); break; + default: launch_gemv(a, stream); break; + } +} - const int n_tiles_m = (ncols + TM - 1) / TM; - const int n_tiles_n = (nrows + TN - 1) / TN; - const int n_tiles = n_tiles_m * n_tiles_n; +template static void launch(const args & a, int ls, dpct::queue_ptr stream); - stream->submit([&](sycl::handler & h) { - if constexpr (S > 1) { - const sycl::nd_range<1> nd{ sycl::range<1>((size_t) n_tiles * S), sycl::range<1>(S) }; - h.parallel_for(nd, [=](sycl::nd_item<1> it) [[intel::sycl_explicit_simd]] { - pq2_xmx_thread(wq, wd, a8, as, dst, K, nrows, ncols, nrows_dst, n_tiles_n, - (int) it.get_group(0), (int) it.get_local_id(0)); - }); - } else { - const size_t global = (size_t) ((n_tiles + PQ2_XMX_WG - 1) / PQ2_XMX_WG) * PQ2_XMX_WG; - const sycl::nd_range<1> nd{ sycl::range<1>(global), sycl::range<1>(PQ2_XMX_WG) }; - h.parallel_for(nd, [=](sycl::nd_item<1> it) [[intel::sycl_explicit_simd]] { - const int tile = (int) it.get_global_id(0); - if (tile < n_tiles) { - pq2_xmx_thread(wq, wd, a8, as, dst, K, nrows, ncols, nrows_dst, n_tiles_n, tile, 0); - } - }); +static sycl::event reorder_pq2_0(const uint8_t * src, uint8_t * dst, int ncols, int nrows, + dpct::queue_ptr stream) { + GGML_ASSERT(ncols % QK_PQ2_0 == 0); + const int nw = ncols / 16; + const int nb = ncols / QK_PQ2_0; + const int np = GGML_PAD(nrows, 16); + uint32_t * qs = (uint32_t *) dst; + sycl::half * d = (sycl::half *) (dst + (size_t) nw * np * sizeof(uint32_t)); + const block_pq2_0 * x = (const block_pq2_0 *) src; + + return stream->parallel_for(sycl::range<2>(nw, np), [=](sycl::item<2> it) { + const int w = it[0]; + const int n = it[1]; + if (n >= nrows) { + qs[(size_t) w * np + n] = 0; + if (w % 8 == 0) { + d[(size_t) (w / 8) * np + n] = 0.0f; + } + return; + } + const block_pq2_0 * blk = x + (size_t) n * nb + w / 8; + // qs is only 2-byte aligned inside the 34-byte block + const uint16_t * q = (const uint16_t *) blk->qs + 2 * (w % 8); + qs[(size_t) w * np + n] = pq2_codes_to_s2(q[0] | ((uint32_t) q[1] << 16)); + if (w % 8 == 0) { + d[(size_t) (w / 8) * np + n] = blk->d; } }); } -// K split for a tile count: the smallest power of two reaching the thread target, at least two blocks per thread -template -static void launch_pq2_xmx_split(const uint32_t * wq, const uint16_t * wd, const uint32_t * a8, const float * as, - float * dst, int K, int nrows, int ncols, int nrows_dst, dpct::queue_ptr stream) { - const int nb = K / QK_PQ2_0; - const int n_tiles = ((ncols + 8 * MR - 1) / (8 * MR)) * ((nrows + 16 * NR - 1) / (16 * NR)); - // 32 token tiles (prefill) carry a large SLM reduction, so they split at most in two and only while short - const int target = MR >= 4 ? PQ2_XMX_PREFILL_TARGET_THREADS : PQ2_XMX_TARGET_THREADS; - const int max_split = MR >= 4 ? 2 : 16; - int split = 1; - while (split < max_split && n_tiles * split < target && nb >= 4 * split) { - split *= 2; +static sycl::event reorder_ptq1_0(const uint8_t * src, uint8_t * dst, int ncols, int nrows, + dpct::queue_ptr stream) { + GGML_ASSERT(ncols % QK_PTQ1_0 == 0); + const int nb = ncols / QK_PTQ1_0; + const int np = GGML_PAD(nrows, 16); + uint32_t * out = (uint32_t *) dst; + const block_ptq1_0 * x = (const block_ptq1_0 *) src; + + return stream->parallel_for(sycl::range<2>((size_t) nb * PTQ1_ROWS, np), [=](sycl::item<2> it) { + const int r = it[0]; + const int n = it[1]; + uint32_t v = 0; + if (n < nrows) { + const block_ptq1_0 * blk = x + (size_t) n * nb + r / PTQ1_ROWS; + const int d = r % PTQ1_ROWS; + if (d < PTQ1_ROWS - 1) { + v = *(const uint32_t *) (blk->qs + 4 * d); // blocks are 28 bytes, so qs is 4-byte aligned + } else { + v = blk->qh[0] | (uint32_t) blk->qh[1] << 8 | (uint32_t) sycl::bit_cast(blk->d) << 16; + } + } + out[(size_t) r * np + n] = v; + }); +} + +// M at or below this uses the GEMV kernel +static constexpr int GEMV_MAX_M = 8; + +// src1 quantized to int8 with one scale per 128 values +struct act_q { + ggml_sycl_pool_alloc aq; + ggml_sycl_pool_alloc sa; + int M, K; + int fmt; // weight format whose K order the activation is in + + act_q(ggml_backend_sycl_context & ctx, const ggml_tensor * src1, int fmt) : aq(ctx.pool()), sa(ctx.pool()), fmt(fmt) { + GGML_ASSERT(src1->type == GGML_TYPE_F32 && src1->nb[0] == sizeof(float)); + GGML_ASSERT(src1->ne[0] % GS == 0); + M = src1->ne[1] * src1->ne[2] * src1->ne[3]; + K = src1->ne[0]; + aq.alloc((size_t) M * K); + sa.alloc((size_t) (K / GS) * ldsa(M)); + + const quant_a q{ (const float *) src1->data, + sa.get(), + aq.get(), + M, + K, + (int) src1->ne[1], + (int) src1->ne[2], + (int64_t) (src1->nb[1] / sizeof(float)), + (int64_t) (src1->nb[2] / sizeof(float)), + (int64_t) (src1->nb[3] / sizeof(float)), + fmt == FMT_PTQ1 }; + ctx.stream()->parallel_for(sycl::nd_range<2>(sycl::range<2>(M, (K / GS) * 16), sycl::range<2>(1, 16)), q); + } +}; + +// out is M x N floats with row stride ldc +static int fmt_of(ggml_type type) { + GGML_ASSERT(type == GGML_TYPE_PQ2_0 || type == GGML_TYPE_PTQ1_0); + return type == GGML_TYPE_PTQ1_0 ? FMT_PTQ1 : FMT_PQ2; +} + +template static void launch(const args & a, int ls, dpct::queue_ptr stream) { + if constexpr (FMT == FMT_PTQ1) { + // decode-bound: tiles that use each decoded block for more rows; with WG_M == 1 the GEMM + // decodes in registers, without the SLM round trip + if (a.M == 1) { + launch_gemv_ls(a, ls, stream); + } else if (a.M <= 4) { + launch_gemv_ls(a, ls, stream); // the 2-row variant is slower + } else if (a.M <= 8 && (int64_t) a.K * a.N < (1 << 24)) { + launch_gemv_ls(a, ls, stream); + } else if (a.M <= 16) { + launch_gemm(a, stream); + } else if (a.M <= 32) { + launch_gemm(a, stream); + } else if (a.M <= 64) { + if (a.K >= 6144) { + launch_gemm(a, stream); + } else { + launch_gemm(a, stream); + } + } else if (a.M <= 128) { + launch_gemm(a, stream); + } else { + // a tall work-group shares each decoded B column block among 8 sub-groups + launch_gemm(a, stream); + } + return; + } + + if (a.M <= GEMV_MAX_M) { + if (a.M == 1) { + launch_gemv_ls(a, ls, stream); + } else if (a.M == 2) { + launch_gemv_ls(a, ls, stream); + } else if (a.M <= 4) { + launch_gemv_ls(a, ls, stream); + } else { + launch_gemv_ls(a, ls, stream); + } + } else if (a.M <= 16) { + // small batches: few row blocks per work-group, so the sub-groups are not idle + launch_gemm(a, stream); + } else if (a.M <= 32) { + launch_gemm(a, stream); + } else if (a.M <= 64) { + launch_gemm(a, stream); + } else if (a.M <= 128) { + launch_gemm(a, stream); + } else { + launch_gemm(a, stream); + } +} + +static void run(ggml_backend_sycl_context & ctx, const ggml_tensor * w, const act_q & q, float * out, int ldc) { + GGML_ASSERT(w->ne[0] == q.K); + const int fmt = fmt_of(w->type); + GGML_ASSERT(fmt == q.fmt); + + const int K = q.K; + const int M = q.M; + const int N = w->ne[1]; + const int NP = GGML_PAD(N, 16); + + // 2D block I/O needs 64-byte aligned surfaces and 16-byte aligned pitches + GGML_ASSERT((uintptr_t) w->data % 64 == 0); + const bool st2d = (uintptr_t) out % 64 == 0 && ldc % 4 == 0; + + args a{ q.aq.ptr, + q.sa.ptr, + (const uint32_t *) w->data, + (const unsigned short *) ((const char *) w->data + (size_t) (K / 16) * NP * sizeof(uint32_t)), + out, + M, + N, + NP, + K, + ldc, + st2d, + &ctx.pool(), + 0 }; + + // split K so the sub-groups fill about one wave of hardware threads (8 per EU); + // a partial second wave costs more than the split saves + const int nblk = NP / 16; + const int threads = ggml_sycl_info().devices[ctx.device].nsm * 16 * 8; // nsm is compute units / 16 + a.threads = threads; + int ls = 4; + for (int l : { 8, 4, 2, 1 }) { + if (nblk * l <= threads) { + ls = l; + break; + } } - switch (split) { - case 1: launch_pq2_xmx(wq, wd, a8, as, dst, K, nrows, ncols, nrows_dst, stream); break; - case 2: launch_pq2_xmx(wq, wd, a8, as, dst, K, nrows, ncols, nrows_dst, stream); break; - case 4: launch_pq2_xmx(wq, wd, a8, as, dst, K, nrows, ncols, nrows_dst, stream); break; - case 8: launch_pq2_xmx(wq, wd, a8, as, dst, K, nrows, ncols, nrows_dst, stream); break; - default: launch_pq2_xmx(wq, wd, a8, as, dst, K, nrows, ncols, nrows_dst, stream); break; + + if (fmt == FMT_PTQ1) { + launch(a, ls, ctx.stream()); + } else { + launch(a, ls, ctx.stream()); } } -} // namespace +} // namespace ggml_sycl_xmx bool ggml_sycl_pq2_xmx_supports_ne0(int64_t ne0) { - return ne0 % QK_PQ2_0 == 0 && (ne0 / QK_PQ2_0) * PQ2_XMX_QS_BYTES >= 64; + return ne0 % QK_PQ2_0 == 0; +} + +size_t ggml_sycl_pq2_xmx_bytes(const ggml_tensor * t) { + return (size_t) GGML_PAD(t->ne[1], 16) * t->nb[1]; } bool ggml_sycl_pq2_xmx_reorder(ggml_tensor * src0, dpct::queue_ptr stream) { GGML_ASSERT((src0->type == GGML_TYPE_PQ2_0 || src0->type == GGML_TYPE_PTQ1_0) && ggml_is_contiguous(src0)); + GGML_ASSERT(src0->ne[2] == 1 && src0->ne[3] == 1); const size_t size = ggml_nbytes(src0); - const size_t nblk = (size_t) ggml_nelements(src0) / QK_PQ2_0; - uint8_t * data = (uint8_t *) src0->data; - - void * tmp = sycl::malloc_device(size, *stream); + void * tmp = sycl::malloc_device(size, *stream); if (!tmp) { - GGML_LOG_WARN("%s: failed to allocate %zu bytes for the PQ2_0 XMX reorder, skipping it\n", __func__, size); + GGML_LOG_WARN("%s: failed to allocate %zu bytes for the XMX reorder, skipping it\n", __func__, size); return false; } - stream->memcpy(tmp, data, size).wait(); - - uint8_t * qs = data; - sycl::half * d = (sycl::half *) (data + nblk * PQ2_XMX_QS_BYTES); - if (src0->type == GGML_TYPE_PQ2_0) { - stream->parallel_for(sycl::range<1>(nblk), [=](sycl::id<1> i) { - const block_pq2_0 * x = (const block_pq2_0 *) tmp + i; -#pragma unroll - for (int j = 0; j < PQ2_XMX_QS_BYTES; ++j) { - qs[i * PQ2_XMX_QS_BYTES + j] = x->qs[j]; - } - d[i] = x->d; - }).wait(); - } else { - // base-3 trits (value -1..1) become PQ2_0 codes (value + 1), four to a byte, lowest first; - // the caller made sure the buffer holds 34 bytes a block - stream->parallel_for(sycl::range<1>(nblk), [=](sycl::id<1> i) { - const block_ptq1_0 * x = (const block_ptq1_0 *) tmp + i; - for (int j = 0; j < PQ2_XMX_QS_BYTES; ++j) { - uint8_t byte = 0; -#pragma unroll - for (int k = 0; k < 4; ++k) { - byte |= (uint8_t) ((ptq1_0_trit(x, 4 * j + k) + 1) << (2 * k)); - } - qs[i * PQ2_XMX_QS_BYTES + j] = byte; - } - d[i] = x->d; - }).wait(); - } - + stream->memcpy(tmp, src0->data, size).wait(); + const auto reorder = src0->type == GGML_TYPE_PQ2_0 ? ggml_sycl_xmx::reorder_pq2_0 : ggml_sycl_xmx::reorder_ptq1_0; + reorder((const uint8_t *) tmp, (uint8_t *) src0->data, (int) src0->ne[0], (int) src0->ne[1], stream).wait(); sycl::free(tmp, *stream); return true; } void ggml_sycl_pq2_xmx_mul_mat(ggml_backend_sycl_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) { - GGML_ASSERT((src0->type == GGML_TYPE_PQ2_0 || src0->type == GGML_TYPE_PTQ1_0) && src0->ne[2] == 1 && - src0->ne[3] == 1); - GGML_ASSERT(src1->type == GGML_TYPE_F32 && src1->nb[0] == sizeof(float)); GGML_ASSERT(dst->type == GGML_TYPE_F32 && ggml_is_contiguous(dst)); - GGML_ASSERT(ggml_sycl_pq2_xmx_supports_ne0(src0->ne[0])); - - const int K = (int) src0->ne[0]; - const int nrows = (int) src0->ne[1]; - const int nb = K / QK_PQ2_0; - const int ne11 = (int) src1->ne[1]; - const int ne12 = (int) src1->ne[2]; - const int ncols = (int) (src1->ne[1] * src1->ne[2] * src1->ne[3]); - - dpct::queue_ptr stream = ctx.stream(); - - // int8 activations (ncols rows of K bytes, 64-byte aligned for 2D loads) and one float scale per 128 values - ggml_sycl_pool_alloc a8_alloc(ctx.pool(), (size_t) ncols * K + 64); - ggml_sycl_pool_alloc as_alloc(ctx.pool(), (size_t) ncols * nb); - int8_t * a8 = (int8_t *) GGML_PAD((uintptr_t) a8_alloc.get(), 64); - float * as = as_alloc.get(); - - { - const char * src1_d = (const char *) src1->data; - const size_t nb11 = src1->nb[1], nb12 = src1->nb[2], nb13 = src1->nb[3]; - // one sub-group per (token, 128-block), 8 values per work-item - stream->parallel_for( - sycl::nd_range<1>(sycl::range<1>((size_t) ncols * nb * 16), sycl::range<1>(16)), - [=](sycl::nd_item<1> it) [[sycl::reqd_sub_group_size(16)]] { - const int grp = (int) it.get_group(0); - const int j = grp / nb; - const int b = grp % nb; - const int l = (int) it.get_local_id(0); - const int i1 = j % ne11; - const int i2 = (j / ne11) % ne12; - const int i3 = j / (ne11 * ne12); - - const float * x = (const float *) (src1_d + i1 * nb11 + i2 * nb12 + i3 * nb13) + b * QK_PQ2_0 + l * 8; - float v[8]; - float amax = 0.0f; -#pragma unroll - for (int i = 0; i < 8; ++i) { - v[i] = x[i]; - amax = sycl::fmax(amax, sycl::fabs(v[i])); - } - amax = sycl::reduce_over_group(it.get_sub_group(), amax, sycl::maximum()); - const float d = amax / 127.0f; - const float id = d != 0.0f ? 1.0f / d : 0.0f; - - sycl::vec q; -#pragma unroll - for (int i = 0; i < 8; ++i) { - q[i] = (int8_t) sycl::round(v[i] * id); - } - *(sycl::vec *) (a8 + (size_t) j * K + b * QK_PQ2_0 + l * 8) = q; - if (l == 0) { - as[(size_t) j * nb + b] = d; - } - }); - } - - const uint32_t * wq = (const uint32_t *) src0->data; - const uint16_t * wd = (const uint16_t *) ((const uint8_t *) src0->data + (size_t) nrows * nb * PQ2_XMX_QS_BYTES); - float * dd = (float *) dst->data; - const int nrows_dst = (int) dst->ne[0]; - - if (ncols <= 8) { - launch_pq2_xmx_split<1, 2>(wq, wd, (const uint32_t *) a8, as, dd, K, nrows, ncols, nrows_dst, stream); - } else if (ncols <= 16) { - launch_pq2_xmx_split<2, 2>(wq, wd, (const uint32_t *) a8, as, dd, K, nrows, ncols, nrows_dst, stream); - } else { - launch_pq2_xmx_split<4, 2>(wq, wd, (const uint32_t *) a8, as, dd, K, nrows, ncols, nrows_dst, stream); - } + const ggml_sycl_xmx::act_q q(ctx, src1, ggml_sycl_xmx::fmt_of(src0->type)); + ggml_sycl_xmx::run(ctx, src0, q, (float *) dst->data, (int) dst->ne[0]); } #else @@ -371,6 +1110,10 @@ bool ggml_sycl_pq2_xmx_supports_ne0(int64_t) { return false; } +size_t ggml_sycl_pq2_xmx_bytes(const ggml_tensor * t) { + return ggml_nbytes(t); +} + bool ggml_sycl_pq2_xmx_reorder(ggml_tensor *, dpct::queue_ptr) { return false; } @@ -378,5 +1121,4 @@ bool ggml_sycl_pq2_xmx_reorder(ggml_tensor *, dpct::queue_ptr) { void ggml_sycl_pq2_xmx_mul_mat(ggml_backend_sycl_context &, const ggml_tensor *, const ggml_tensor *, ggml_tensor *) { GGML_ABORT("PQ2_0 XMX path is not built in"); } - #endif // __INTEL_LLVM_COMPILER && !GGML_SYCL_NO_PQ2_XMX diff --git a/ggml/src/ggml-sycl/pq2_xmx.hpp b/ggml/src/ggml-sycl/pq2_xmx.hpp index 4ad3c7493f0b..f2be23a5eaed 100644 --- a/ggml/src/ggml-sycl/pq2_xmx.hpp +++ b/ggml/src/ggml-sycl/pq2_xmx.hpp @@ -2,16 +2,22 @@ #include "common.hpp" -// PQ2_0 in the XMX layout: the weight tensor is rewritten in place, once, into a plane of 32-byte qs blocks -// (row pitch nb * 32 bytes, so every row is 2D-block-load aligned) followed by a plane of fp16 block scales. -// PTQ1_0 weights take the same layout: their base-3 trits are expanded to PQ2_0 codes, which needs 34 bytes a -// block instead of 28, so the buffer type reserves that room for them on devices that use this path. -// Activations are quantized to int8 with one float scale per 128 values, so the four DPAS of a PQ2_0 block -// accumulate in integers before a single float rescale. +// PQ2_0 / PTQ1_0 x int8 on 16-wide DPAS devices with native s2 x s8 DPAS. The weight tensor is rewritten in +// place, once, into an XMX layout whose rows are padded to NP = ne[1] rounded up to 16 (zero columns), so the +// buffer must hold ggml_sycl_pq2_xmx_bytes(): +// PQ2_0: uint32 qs[K/16][NP] as s2 codes, then half d[K/128][NP]. Codes are ternary only: PQ2_0 code 3 (+2) +// has no s2 value. +// PTQ1_0: uint32 [K/128][7][NP]: the 24 qs bytes, then qh[0] | qh[1] << 8 | d << 16. The trits are decoded +// in the kernels, so the weights stay at 1.75 bits. +// Activations are quantized to int8 with one float scale per 128 values, so the four DPAS of a block accumulate +// in integers before a single float rescale. -// ne[0] of a PQ2_0 weight the XMX path accepts: a 2D surface needs a row of at least 64 bytes +// ne[0] of a weight the XMX path accepts bool ggml_sycl_pq2_xmx_supports_ne0(int64_t ne0); +// bytes a 2D weight takes in the XMX layout +size_t ggml_sycl_pq2_xmx_bytes(const ggml_tensor * t); + // rewrite src0 (PQ2_0 or PTQ1_0, AoS blocks) into the XMX layout in place bool ggml_sycl_pq2_xmx_reorder(ggml_tensor * src0, dpct::queue_ptr stream);