diff --git a/ggml/src/ggml-sycl/common.hpp b/ggml/src/ggml-sycl/common.hpp index 8f2106d83088..98d84e2735e9 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 { @@ -440,12 +440,17 @@ struct ggml_backend_sycl_context { std::unique_ptr host_pools[GGML_SYCL_MAX_DEVICES]; + // XMX activations shared across mat-muls are freed out of order, which the VMM pool does not allow + std::unique_ptr xmx_act_pools[GGML_SYCL_MAX_DEVICES]; + std::vector mmid_row_mapping_host; static std::unique_ptr new_pool_for_device(queue_ptr qptr, int device); static std::unique_ptr new_pool_for_host(queue_ptr qptr, int device); + static std::unique_ptr new_unordered_pool_for_device(queue_ptr qptr, int device); + static std::unique_ptr new_fattn_kv_buffers(queue_ptr qptr, int device); ggml_sycl_pool & pool(int device) { @@ -459,6 +464,13 @@ struct ggml_backend_sycl_context { return pool(device); } + ggml_sycl_pool & xmx_act_pool() { + if (xmx_act_pools[device] == nullptr) { + xmx_act_pools[device] = new_unordered_pool_for_device(stream(device, 0), device); + } + return *xmx_act_pools[device]; + } + ggml_sycl_fattn_kv_buffers & fattn_buffers(int device) { if (fattn_bufs[device] == nullptr) { fattn_bufs[device] = new_fattn_kv_buffers(stream(device, 0), device); diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp index 97b5f35659f7..f4086503a684 100644 --- a/ggml/src/ggml-sycl/ggml-sycl.cpp +++ b/ggml/src/ggml-sycl/ggml-sycl.cpp @@ -24,6 +24,8 @@ #include #include #include +#include +#include #include #include #include @@ -983,16 +985,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 +1004,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; @@ -1932,6 +1930,10 @@ std::unique_ptr ggml_backend_sycl_context::new_pool_for_device(q } +std::unique_ptr ggml_backend_sycl_context::new_unordered_pool_for_device(queue_ptr qptr, int device) { + return std::unique_ptr(new ggml_sycl_pool_leg(qptr, device)); +} + std::unique_ptr ggml_backend_sycl_context::new_fattn_kv_buffers(queue_ptr qptr, int device) { return std::unique_ptr(new ggml_sycl_fattn_kv_buffers(qptr, device)); } @@ -3808,7 +3810,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) { @@ -4582,8 +4584,9 @@ static bool can_use_mul_mat_vec_q(const ggml_tensor * src0, const ggml_tensor * // PQ2_0/PTQ1_0 weights on 16-wide DPAS devices are rewritten into the XMX layout on first use. From then on every // mul_mat on them has to take that path, so the layout flag alone decides once it is set. -static bool ggml_sycl_pq2_xmx_use(ggml_backend_sycl_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1, - const ggml_tensor * dst) { +// ggml_sycl_pq2_xmx_use() without the rewrite: whether this mul_mat would take the XMX path +static bool ggml_sycl_pq2_xmx_eligible(ggml_backend_sycl_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1, + const ggml_tensor * dst) { if (src0->type != GGML_TYPE_PQ2_0 && src0->type != GGML_TYPE_PTQ1_0) { return false; } @@ -4612,11 +4615,19 @@ 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 + return ggml_backend_buft_get_alloc_size(src0->buffer->buft, src0) >= ggml_sycl_pq2_xmx_bytes(src0); +} + +static bool ggml_sycl_pq2_xmx_use(ggml_backend_sycl_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1, + const ggml_tensor * dst) { + if (!ggml_sycl_pq2_xmx_eligible(ctx, src0, src1, dst)) { return false; } + ggml_tensor_extra_gpu * extra = static_cast(src0->extra); + if (extra->optimized_feature.xmx_pq2) { + return true; + } if (!ggml_sycl_pq2_xmx_reorder(const_cast(src0), ctx.stream())) { return false; } @@ -4624,6 +4635,316 @@ static bool ggml_sycl_pq2_xmx_use(ggml_backend_sycl_context & ctx, const ggml_te return true; } +static bool ggml_sycl_is_view_or_noop(const ggml_tensor * t); + +// Quantized XMX activations shared within one graph compute. An entry is keyed by the activation +// tensor (a reshape or view of all of it resolves to the same key) and lives until its last XMX +// consumer ran. Holding it also lets a mat-mul run later than its graph position. +struct ggml_sycl_xmx_graph_state { + struct act_entry { + ggml_sycl_pq2_xmx_act * act; + int remaining; + }; + + // keyed by (activation, weight type): PQ2_0 and PTQ1_0 read the activation in different K orders + std::map, act_entry> acts; + std::unordered_map deferred; // GLU node -> its gate mat-mul + + ~ggml_sycl_xmx_graph_state() { + for (auto & e : acts) { + ggml_sycl_pq2_xmx_act_free(e.second.act); + } + } +}; + +// consumers sharing a plain activation sit close together (q/k/v, gate/up); later ones quantize again +static constexpr int GGML_SYCL_XMX_SHARE_WINDOW = 64; + +static const ggml_tensor * ggml_sycl_xmx_act_key(const ggml_tensor * b) { + const ggml_tensor * r = b->view_src; + return r && b->data == r->data && ggml_nelements(b) == ggml_nelements(r) ? r : b; +} + +static bool ggml_sycl_is_alias_op(const ggml_tensor * n) { + return n->op == GGML_OP_RESHAPE || n->op == GGML_OP_VIEW || n->op == GGML_OP_PERMUTE || n->op == GGML_OP_TRANSPOSE; +} + +// a mat-mul that the graph loop runs on the XMX kernels from a shared quantized activation +static bool ggml_sycl_is_xmx_mul_mat(ggml_backend_sycl_context & ctx, const ggml_tensor * n) { + if (n->op != GGML_OP_MUL_MAT || !ggml_sycl_pq2_xmx_eligible(ctx, n->src[0], n->src[1], n)) { + return false; + } + const ggml_tensor * b = n->src[1]; + return b->type == GGML_TYPE_F32 && ggml_is_contiguous(b) && b->ne[0] == n->src[0]->ne[0] && + n->type == GGML_TYPE_F32 && ggml_is_contiguous(n) && ggml_nelements(n) == n->ne[0] * ggml_nrows(b); +} + +static int ggml_sycl_use_count(const ggml_cgraph * cgraph, const ggml_tensor * t) { + const size_t pos = ggml_hash_find(&cgraph->visited_hash_set, t); + if (pos == GGML_HASHSET_FULL || !ggml_bitset_get(cgraph->visited_hash_set.used, pos)) { + return -1; + } + return cgraph->use_counts[pos]; +} + +// Count the users of t (seen through alias ops) from node `from` on, stopping once the graph's +// use counts are all accounted for. Returns -1 if some are missing (e.g. in another split). +// n_xmx gets the users that are XMX mat-muls reading t as src1, and wtype their weight type +// (GGML_TYPE_COUNT if they differ). +static int ggml_sycl_xmx_users(ggml_backend_sycl_context & ctx, const ggml_cgraph * cgraph, int from, + const ggml_tensor * t, int & n_xmx, ggml_type & wtype) { + wtype = GGML_TYPE_COUNT; + int pending = ggml_sycl_use_count(cgraph, t); + if (pending < 0) { + return -1; + } + std::vector aliases = { t }; + int total = 0; + n_xmx = 0; + for (int j = from; j < cgraph->n_nodes && pending > 0; ++j) { + const ggml_tensor * n = cgraph->nodes[j]; + for (int s = 0; s < GGML_MAX_SRC && n->src[s]; ++s) { + if (std::find(aliases.begin(), aliases.end(), n->src[s]) == aliases.end()) { + continue; + } + pending--; + if (ggml_sycl_is_alias_op(n)) { + const int uc = ggml_sycl_use_count(cgraph, n); + if (uc < 0) { + return -1; + } + aliases.push_back(n); + pending += uc; + continue; + } + total++; + if (s == 1 && ggml_sycl_is_xmx_mul_mat(ctx, n) && ggml_sycl_xmx_act_key(n->src[1]) == t) { + const ggml_type ty = n->src[0]->type; + wtype = n_xmx == 0 || wtype == ty ? ty : GGML_TYPE_COUNT; + n_xmx++; + } + } + } + return pending == 0 ? total : -1; +} + +static ggml_sycl_pq2_xmx_act * ggml_sycl_xmx_acquire(ggml_backend_sycl_context & ctx, ggml_sycl_xmx_graph_state & st, + const ggml_cgraph * cgraph, int node_idx) { + const ggml_tensor * mm = cgraph->nodes[node_idx]; + const ggml_type wtype = mm->src[0]->type; + const auto key = std::make_pair(ggml_sycl_xmx_act_key(mm->src[1]), wtype); + const auto it = st.acts.find(key); + if (it != st.acts.end()) { + return it->second.act; + } + int users = 0; + const int end = std::min(cgraph->n_nodes, node_idx + GGML_SYCL_XMX_SHARE_WINDOW); + for (int j = node_idx; j < end; ++j) { + const ggml_tensor * n = cgraph->nodes[j]; + if (ggml_sycl_is_xmx_mul_mat(ctx, n) && n->src[0]->type == wtype && + ggml_sycl_xmx_act_key(n->src[1]) == key.first) { + users++; + } + } + ggml_sycl_pq2_xmx_act * act = ggml_sycl_pq2_xmx_act_quantize(ctx, mm->src[1], nullptr, wtype); + st.acts[key] = { act, users }; + return act; +} + +static void ggml_sycl_xmx_release(ggml_sycl_xmx_graph_state & st, const ggml_tensor * mm) { + const auto it = st.acts.find(std::make_pair(ggml_sycl_xmx_act_key(mm->src[1]), mm->src[0]->type)); + GGML_ASSERT(it != st.acts.end()); + if (--it->second.remaining <= 0) { + ggml_sycl_pq2_xmx_act_free(it->second.act); + st.acts.erase(it); + } +} + +// Skip RESHAPE/VIEW nodes after idx that alias t contiguously. Returns the last alias (or t) +// and sets idx to the first other node. +static const ggml_tensor * ggml_sycl_skip_aliases(const ggml_cgraph * cgraph, const ggml_tensor * t, int & idx) { + const ggml_tensor * cur = t; + for (; idx < cgraph->n_nodes; ++idx) { + const ggml_tensor * n = cgraph->nodes[idx]; + if ((n->op != GGML_OP_RESHAPE && n->op != GGML_OP_VIEW) || n->src[0] != cur || n->data != t->data || + !ggml_is_contiguous(n) || (n->flags & GGML_TENSOR_FLAG_OUTPUT)) { + break; + } + cur = n; + } + return cur; +} + +// Run the XMX mat-mul at node_idx and fuse what follows it where possible: +// {mul_mat(gate), mul_mat(up), SWIGLU}, a later SWIGLU whose gate it is (run deferred, at the GLU), +// and {mul_mat, reshape/view..., ADD} residuals. Returns the number of nodes consumed after +// node_idx, or -1 if the node is not handled here. +static int ggml_sycl_xmx_mul_mat_node(ggml_backend_sycl_context & ctx, ggml_sycl_xmx_graph_state & st, + ggml_cgraph * cgraph, int node_idx) { + const ggml_tensor * mm = cgraph->nodes[node_idx]; + if (!ggml_sycl_is_xmx_mul_mat(ctx, mm) || !ggml_sycl_pq2_xmx_use(ctx, mm->src[0], mm->src[1], mm)) { + return -1; + } + const int N = mm->ne[0]; + + if (g_ggml_sycl_enable_fusion && + ggml_can_fuse_subgraph(cgraph, node_idx, { GGML_OP_MUL_MAT, GGML_OP_MUL_MAT, GGML_OP_GLU }, { node_idx + 2 })) { + ggml_tensor * glu = cgraph->nodes[node_idx + 2]; + const ggml_tensor * gate = glu->src[0]; + const ggml_tensor * up = glu->src[1]; + const bool pair = (gate == mm && up == cgraph->nodes[node_idx + 1]) || + (up == mm && gate == cgraph->nodes[node_idx + 1]); + if (pair && ggml_get_glu_op(glu) == GGML_GLU_OP_SWIGLU && !ggml_get_op_params_i32(glu, 1) /* swapped */ && + ggml_sycl_is_xmx_mul_mat(ctx, gate) && ggml_sycl_is_xmx_mul_mat(ctx, up) && + ggml_sycl_xmx_act_key(gate->src[1]) == ggml_sycl_xmx_act_key(up->src[1]) && + gate->src[0]->type == up->src[0]->type && + ggml_are_same_shape(gate, up) && ggml_are_same_shape(glu, gate) && glu->type == GGML_TYPE_F32 && + ggml_is_contiguous(glu)) { + scope_op_debug_print scope_dbg_print(__func__, glu, /*num_src=*/2, " : xmx gate + up + swiglu"); + GGML_ASSERT(ggml_sycl_pq2_xmx_use(ctx, gate->src[0], gate->src[1], gate) && + ggml_sycl_pq2_xmx_use(ctx, up->src[0], up->src[1], up)); + const ggml_sycl_pq2_xmx_act * act = ggml_sycl_xmx_acquire(ctx, st, cgraph, node_idx); + ggml_sycl_pool_alloc up_out(ctx.pool(), ggml_nelements(up)); + ggml_sycl_pq2_xmx_mul_mat_act(ctx, up->src[0], act, up_out.get(), N, GGML_SYCL_XMX_EPI_NONE, nullptr); + ggml_sycl_pq2_xmx_mul_mat_act(ctx, gate->src[0], act, (float *) glu->data, N, + GGML_SYCL_XMX_EPI_SWIGLU, up_out.get()); + ggml_sycl_xmx_release(st, gate); + ggml_sycl_xmx_release(st, up); + return 2; + } + } + + // single user, reached through aliases: a SWIGLU that takes this as its gate, or a residual ADD + const ggml_tensor * user = nullptr; + const ggml_tensor * cur = mm; + int user_idx = -1; + if (g_ggml_sycl_enable_fusion && ggml_node_get_use_count(cgraph, node_idx) == 1 && + !(mm->flags & GGML_TENSOR_FLAG_OUTPUT)) { + const int end = std::min(cgraph->n_nodes, node_idx + 2 * GGML_SYCL_XMX_SHARE_WINDOW); + for (int j = node_idx + 1; j < end && !user; ++j) { + const ggml_tensor * n = cgraph->nodes[j]; + const bool uses_cur = n->src[0] == cur || n->src[1] == cur; + if (!uses_cur) { + continue; + } + if ((n->op == GGML_OP_RESHAPE || n->op == GGML_OP_VIEW) && n->src[0] == cur && n->data == mm->data && + ggml_is_contiguous(n) && ggml_node_get_use_count(cgraph, j) == 1 && + !(n->flags & GGML_TENSOR_FLAG_OUTPUT)) { + cur = n; + continue; + } + user = n; + user_idx = j; + } + } + + if (user && user->op == GGML_OP_GLU && ggml_get_glu_op(user) == GGML_GLU_OP_SWIGLU && user->src[0] == cur && + !ggml_get_op_params_i32(user, 1) /* swapped */ && user->src[1] && user->src[1] != cur && + user->src[1]->type == GGML_TYPE_F32 && ggml_is_contiguous(user->src[1]) && + ggml_nelements(user->src[1]) == ggml_nelements(mm) && user->type == GGML_TYPE_F32 && + ggml_is_contiguous(user) && ggml_nelements(user) == ggml_nelements(mm)) { + // the multiplier is computed later; keep the quantized activation and run at the GLU + ggml_sycl_xmx_acquire(ctx, st, cgraph, node_idx); + st.deferred[user] = mm; + return 0; + } + + if (user && user->op == GGML_OP_ADD && (user->flags & GGML_TENSOR_FLAG_COMPUTE)) { + const ggml_tensor * other = user->src[0] == cur ? user->src[1] : (user->src[1] == cur ? user->src[0] : nullptr); + if (other && other->type == GGML_TYPE_F32 && user->type == GGML_TYPE_F32 && ggml_is_contiguous(other) && + ggml_is_contiguous(user) && ggml_are_same_shape(user->src[0], user->src[1]) && + ggml_nelements(user) == ggml_nelements(mm)) { + bool only_aliases = true; // nothing but the alias chain may run between the two + for (int j = node_idx + 1; j < user_idx; ++j) { + only_aliases = only_aliases && ggml_sycl_is_view_or_noop(cgraph->nodes[j]); + } + if (only_aliases) { + scope_op_debug_print scope_dbg_print(__func__, user, /*num_src=*/2, " : xmx mul_mat + add"); + const ggml_sycl_pq2_xmx_act * act = ggml_sycl_xmx_acquire(ctx, st, cgraph, node_idx); + ggml_sycl_pq2_xmx_mul_mat_act(ctx, mm->src[0], act, (float *) user->data, N, GGML_SYCL_XMX_EPI_ADD, + (const float *) other->data); + ggml_sycl_xmx_release(st, mm); + return user_idx - node_idx; + } + } + } + + scope_op_debug_print scope_dbg_print(__func__, mm, /*num_src=*/2, " : xmx"); + const ggml_sycl_pq2_xmx_act * act = ggml_sycl_xmx_acquire(ctx, st, cgraph, node_idx); + ggml_sycl_pq2_xmx_mul_mat_act(ctx, mm->src[0], act, (float *) mm->data, N, GGML_SYCL_XMX_EPI_NONE, nullptr); + ggml_sycl_xmx_release(st, mm); + return 0; +} + +// the GLU of a deferred gate mat-mul: run the mat-mul with the SWIGLU epilogue into the GLU output +static bool ggml_sycl_xmx_deferred_glu(ggml_backend_sycl_context & ctx, ggml_sycl_xmx_graph_state & st, + ggml_tensor * glu) { + const auto it = st.deferred.find(glu); + if (it == st.deferred.end()) { + return false; + } + const ggml_tensor * mm = it->second; + st.deferred.erase(it); + scope_op_debug_print scope_dbg_print(__func__, glu, /*num_src=*/2, " : xmx mul_mat + swiglu (deferred)"); + const auto a = st.acts.find(std::make_pair(ggml_sycl_xmx_act_key(mm->src[1]), mm->src[0]->type)); + GGML_ASSERT(a != st.acts.end()); + ggml_sycl_pq2_xmx_mul_mat_act(ctx, mm->src[0], a->second.act, (float *) glu->data, mm->ne[0], + GGML_SYCL_XMX_EPI_SWIGLU, (const float *) glu->src[1]->data); + ggml_sycl_xmx_release(st, mm); + return true; +} + +// Hadamard-folded PQ2_0 input: {x * signs, FWHT_1024 (hinted mul_mat)}. If every user of the FWHT +// output is an XMX mat-mul, the sign flip and the FWHT go into its activation quantizer and the +// output is never written; otherwise one kernel writes it. Returns the nodes consumed after node_idx. +static int ggml_sycl_hadamard_xmx_fused(ggml_backend_sycl_context & ctx, ggml_sycl_xmx_graph_state & st, + ggml_cgraph * cgraph, int node_idx) { + if (!g_ggml_sycl_enable_fusion) { + return 0; + } + const ggml_tensor * mul = cgraph->nodes[node_idx]; + if (mul->op != GGML_OP_MUL || mul->type != GGML_TYPE_F32 || !ggml_is_contiguous(mul) || + ggml_node_get_use_count(cgraph, node_idx) != 1 || (mul->flags & GGML_TENSOR_FLAG_OUTPUT)) { + return 0; + } + const ggml_tensor * x = mul->src[0]; + const ggml_tensor * signs = mul->src[1]; + if (x->type != GGML_TYPE_F32 || signs->type != GGML_TYPE_F32 || !ggml_are_same_shape(x, mul) || + !ggml_is_contiguous(x) || !ggml_is_contiguous(signs) || ggml_nelements(signs) != x->ne[0] || + x->ne[0] % 1024 != 0) { + return 0; + } + + int j = node_idx + 1; + const ggml_tensor * cur = ggml_sycl_skip_aliases(cgraph, mul, j); + if (j >= cgraph->n_nodes) { + return 0; + } + ggml_tensor * had = cgraph->nodes[j]; + if (had->op != GGML_OP_MUL_MAT || ggml_get_op_params_i32(had, 1) != GGML_HINT_SRC0_IS_HADAMARD || + had->src[1] != cur || had->src[0]->ne[0] != 1024 || had->src[0]->ne[1] != 1024 || + had->type != GGML_TYPE_F32 || !ggml_is_contiguous(had) || ggml_nelements(had) != ggml_nelements(mul) || + (had->flags & GGML_TENSOR_FLAG_OUTPUT)) { + return 0; + } + + const float * sd = (const float *) signs->data; + int n_xmx = 0; + ggml_type wtype = GGML_TYPE_COUNT; + const int total = ggml_sycl_xmx_users(ctx, cgraph, j + 1, had, n_xmx, wtype); + if (total > 0 && n_xmx == total && wtype != GGML_TYPE_COUNT && !st.acts.count({ had, wtype })) { + scope_op_debug_print scope_dbg_print(__func__, had, /*num_src=*/2, " : hadamard into xmx quantizer"); + st.acts[{ had, wtype }] = { ggml_sycl_pq2_xmx_act_quantize(ctx, x, sd, wtype), n_xmx }; + return j - node_idx; + } + + scope_op_debug_print scope_dbg_print(__func__, had, /*num_src=*/2, " : hadamard signs + fwht"); + ggml_sycl_pq2_xmx_hadamard_fwht(ctx, x, sd, (float *) had->data); + return j - node_idx; +} + + + static void ggml_sycl_mul_mat(ggml_backend_sycl_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) { scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/2); @@ -5783,6 +6104,8 @@ static int ggml_sycl_try_gdn_cache_fusion(const ggml_cgraph * cgraph, int node_i static void ggml_backend_sycl_graph_compute_impl(ggml_backend_sycl_context * sycl_ctx, ggml_cgraph * cgraph) { ggml_sycl_set_main_device(sycl_ctx->device); + ggml_sycl_xmx_graph_state xmx_state; + for (int i = 0; i < cgraph->n_nodes; i++) { ggml_tensor * node = cgraph->nodes[i]; if (ggml_sycl_is_view_or_noop(node)) { @@ -5828,6 +6151,24 @@ static void ggml_backend_sycl_graph_compute_impl(ggml_backend_sycl_context * syc continue; } + if (node->op == GGML_OP_GLU && ggml_sycl_xmx_deferred_glu(*sycl_ctx, xmx_state, node)) { + continue; + } + if (node->op == GGML_OP_MUL) { + const int had_skip = ggml_sycl_hadamard_xmx_fused(*sycl_ctx, xmx_state, cgraph, i); + if (had_skip > 0) { + i += had_skip; + continue; + } + } + if (node->op == GGML_OP_MUL_MAT) { + const int xmx_skip = ggml_sycl_xmx_mul_mat_node(*sycl_ctx, xmx_state, cgraph, i); + if (xmx_skip >= 0) { + i += xmx_skip; + continue; + } + } + if (node->op == GGML_OP_MUL_MAT && ggml_sycl_mul_mat_glu_mmvq_fused(*sycl_ctx, cgraph, i)) { i += 2; continue; diff --git a/ggml/src/ggml-sycl/pq2_xmx.cpp b/ggml/src/ggml-sycl/pq2_xmx.cpp index 20fbf60f4ae9..ddf43512a46a 100644 --- a/ggml/src/ggml-sycl/pq2_xmx.cpp +++ b/ggml/src/ggml-sycl/pq2_xmx.cpp @@ -1,368 +1,1335 @@ +// +// 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; +} + +template inline void static_for_impl(F && f, std::integer_sequence) { + (f(std::integral_constant{}), ...); +} -namespace { +template inline void static_for(F && f) { + static_for_impl(f, std::make_integer_sequence{}); +} -namespace esimd = sycl::ext::intel::esimd; -namespace xmx = sycl::ext::intel::esimd::xmx; +// 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; +} -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 +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; +} -// 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; +// 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 -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"); +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; +} - 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); +// 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; +} - // 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; +// 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; + } + if (p < 120) { + return 80 + ((p - 80) % 5) * 8 + (p - 80) / 5; + } + return 120 + ((p - 120) % 4) * 2 + (p - 120) / 4; +} - // scale gathers clamp to a valid row/token; the clamped lanes are never stored - simd d_off[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 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)); + for (int m = 0; m < 16; ++m) { + e[m] = lut[(r[m / 4] >> (8 * (m % 4))) & 0xFF]; } - 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)); + 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; +} + +// fused epilogues on the fp32 result, numbered as TernSYCL postops. other has the layout of C. +enum { EPI_NONE = 0, EPI_SWIGLU = 1, EPI_ADD = 2 }; + +inline float epilogue(float v, const float * other, size_t i, int postop) { + if (postop == EPI_SWIGLU) { + return v / (1.0f + sycl::exp(-v)) * other[i]; + } + if (postop == EPI_ADD) { + return v + other[i]; + } + return v; +} + +// 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 + + 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; - simd acc[MR][NR]; + float v[8]; + float mx = 0.0f; #pragma unroll - for (int s = 0; s < MR; ++s) { + 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 g = 0; g < NR; ++g) { - acc[s][g] = 0.0f; + 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 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> }; } +}; + +// quant_a with a sign flip and a 1024-wide normalized Walsh-Hadamard transform in front, as the +// Hadamard-folded weights expect: Aq = quant(FWHT(A * signs)). One work-group per (row, 1024-block). +// The butterflies follow ggml_sycl_op_fwht (fwht_kernel_wide, NT = 256). +struct quant_a_had { + static constexpr int HN = 1024; + static constexpr int NT = 256; + static constexpr int EL = HN / NT; - 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]; + const float * A; + const float * signs; + float * SA; + int8_t * Aq; + float * out; // if set, write the transformed rows (contiguous [K, M]) instead of quantizing + int M, K, ne11, ne12; + int64_t s11, s12, s13; + int perm; // quantize each 128-group in PTQ1_0 decode order + + void operator()(sycl::nd_item<2> it) const { + float * smem = *sycl::ext::oneapi::group_local_memory_for_overwrite(it.get_group()); + const auto sg = it.get_sub_group(); + const int lane = sg.get_local_linear_id(); + const int tid = it.get_local_id(1); + const int b = it.get_group(1); + const int m = it.get_group(0); + + 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 + b * HN; + const float * sg_ = signs + b * HN; + + float reg[EL]; #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); + for (int i = 0; i < EL; ++i) { + reg[i] = a[i * NT + tid] * sg_[i * NT + tid] * (1.0f / 32.0f); // 1 / sqrt(1024) } - + // butterflies inside the sub-group #pragma unroll - for (int s = 0; s < MR; ++s) { - const int yrow = m0 + 8 * s; - - simd ci[NR]; + for (int h = 1; h < 16; h *= 2) { +#pragma unroll + for (int j = 0; j < EL; ++j) { + const float v = reg[j]; + const float v2 = sycl::permute_group_by_xor(sg, v, h); + reg[j] = (lane & h) == 0 ? v + v2 : v2 - v; + } + } + // across sub-groups, through local memory + for (int h = 16; h < NT; h *= 2) { +#pragma unroll + for (int j = 0; j < EL; ++j) { + smem[j * NT + tid] = reg[j]; + } + sycl::group_barrier(it.get_group()); #pragma unroll - for (int g = 0; g < NR; ++g) { - ci[g] = 0; + for (int j = 0; j < EL; ++j) { + const float v = reg[j]; + const float v2 = smem[j * NT + (tid ^ h)]; + reg[j] = (tid & h) == 0 ? v + v2 : v2 - v; } + sycl::group_barrier(it.get_group()); + } + // across registers +#pragma unroll + for (int h = NT; h < HN; h *= 2) { + const int step = h / NT; #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 j = 0; j < EL; j += 2 * step) { +#pragma unroll + for (int k = 0; k < step; ++k) { + const float x = reg[j + k]; + const float y = reg[j + k + step]; + reg[j + k] = x + y; + reg[j + k + step] = x - y; } } + } + if (out) { + float * o = out + (size_t) m * K + b * HN; +#pragma unroll + for (int j = 0; j < EL; ++j) { + o[j * NT + tid] = reg[j]; + } + return; + } +#pragma unroll + for (int j = 0; j < EL; ++j) { + smem[j * NT + tid] = reg[j]; + } + sycl::group_barrier(it.get_group()); - const simd da = gather(as, s_off[s] + (uint32_t) (b * sizeof(float))); + // quantize as quant_a: sub-group g < 8 takes 128-group g of this block + const int g = tid / 16; + if (g >= HN / GS) { + return; + } + float v[8]; + float mx = 0.0f; #pragma unroll - for (int g = 0; g < NR; ++g) { + for (int i = 0; i < 8; ++i) { + v[i] = smem[g * GS + 8 * lane + 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); + const int gg = b * (HN / GS) + g; + if (perm) { #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 i = 0; i < 8; ++i) { + v[i] = smem[g * GS + ptq1_k_orig(8 * lane + i)]; } } + if (lane == 0) { + SA[(size_t) gg * ldsa(M) + m] = s; + } + uint64_t q = 0; +#pragma unroll + 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 + gg * GS + 8 * lane) = q; } - if constexpr (S > 1) { - // every thread parks its partial tile in SLM; thread 0 sums them and stores + auto get(syclex::properties_tag) const { + return syclex::properties{ syclex::sub_group_size<16>, syclex::work_group_size<1, NT> }; + } +}; + +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; + const float * other; + int M, N, NP, K, ldc, postop; + + 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 s = 0; s < MR; ++s) { + 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 g = 0; g < NR; ++g) { + 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 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) + (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); + } + }; + + 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 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 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); } } } - barrier(); - if (ks != 0) { + + 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 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 r = 0; r < SGM; ++r) { + set_el(acc, r, el(acc, r) + src[r * 16 + lane]); + } + } + } + + 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) { + const size_t i = (size_t) (m0 + r) * ldc + n0 + lane; + C[i] = epilogue(el(acc, r), other, i, postop); + } + } + } + + 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; + const float * other; + int M, N, NP, K, ldc, postop; + 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 g = 0; g < NR; ++g) { + 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 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 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 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; + // an epilogue is applied per element on the way out, so acc stays whole for the 2D store + const int epi = ks > 1 ? EPI_NONE : postop; + const bool s2d = ks > 1 || (st2d && epi == EPI_NONE); + 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] = epilogue(acc[i][j][r], other, (size_t) m * ldo + n, epi); + } + } + } } } } + + 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; + const float * other; + int M, N, NP, K, ldc, postop; + 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.other, a.M, a.N, a.NP, a.K, a.ldc, a.postop }); } -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_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.other, a.M, a.N, a.NP, a.K, a.ldc, a.postop, 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; + const float * other = a.other; + const int postop = a.postop; + 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]; + } + const size_t i = (size_t) m * ldc + n; + C[i] = epilogue(v, other, i, postop); + }); + } +} - 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_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; + } +} - 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); - } - }); +template static void launch(const args & a, int ls, dpct::queue_ptr stream); + +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; + } + }); +} + +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; }); } -// 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; +// 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; shared by every weight that reads src1 +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, ggml_sycl_pool & pool, const ggml_tensor * src1, const float * had_signs, + int fmt) : + aq(pool), + sa(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)); + + if (had_signs) { + GGML_ASSERT(K % quant_a_had::HN == 0); + const quant_a_had q{ (const float *) src1->data, + had_signs, + sa.get(), + aq.get(), + nullptr, + 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 / quant_a_had::HN) * quant_a_had::NT), + sycl::range<2>(1, quant_a_had::NT)), + q); + return; + } + + 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 (and other) are 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, + int postop, const float * other) { + 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, + other, + M, + N, + NP, + K, + ldc, + postop, + 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 + +static_assert((int) GGML_SYCL_XMX_EPI_NONE == ggml_sycl_xmx::EPI_NONE && + (int) GGML_SYCL_XMX_EPI_SWIGLU == ggml_sycl_xmx::EPI_SWIGLU && + (int) GGML_SYCL_XMX_EPI_ADD == ggml_sycl_xmx::EPI_ADD, "epilogue numbering"); + +struct ggml_sycl_pq2_xmx_act { + ggml_sycl_xmx::act_q q; + + ggml_sycl_pq2_xmx_act(ggml_backend_sycl_context & ctx, const ggml_tensor * x, const float * had_signs, + ggml_type wtype) : + q(ctx, ctx.xmx_act_pool(), x, had_signs, ggml_sycl_xmx::fmt_of(wtype)) {} +}; + +ggml_sycl_pq2_xmx_act * ggml_sycl_pq2_xmx_act_quantize(ggml_backend_sycl_context & ctx, const ggml_tensor * x, + const float * had_signs, ggml_type wtype) { + return new ggml_sycl_pq2_xmx_act(ctx, x, had_signs, wtype); +} + +void ggml_sycl_pq2_xmx_act_free(ggml_sycl_pq2_xmx_act * act) { + delete act; +} + +void ggml_sycl_pq2_xmx_mul_mat_act(ggml_backend_sycl_context & ctx, const ggml_tensor * w, + const ggml_sycl_pq2_xmx_act * act, float * dst, int ldc, int epi, + const float * other) { + ggml_sycl_xmx::run(ctx, w, act->q, dst, ldc, epi, other); +} 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]; + const ggml_sycl_xmx::act_q q(ctx, ctx.pool(), src1, nullptr, ggml_sycl_xmx::fmt_of(src0->type)); + ggml_sycl_xmx::run(ctx, src0, q, (float *) dst->data, (int) dst->ne[0], ggml_sycl_xmx::EPI_NONE, nullptr); +} - 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); - } +void ggml_sycl_pq2_xmx_hadamard_fwht(ggml_backend_sycl_context & ctx, const ggml_tensor * x, const float * signs, + float * dst) { + using ggml_sycl_xmx::quant_a_had; + GGML_ASSERT(x->type == GGML_TYPE_F32 && x->nb[0] == sizeof(float) && x->ne[0] % quant_a_had::HN == 0); + const int M = x->ne[1] * x->ne[2] * x->ne[3]; + const int K = x->ne[0]; + const quant_a_had q{ (const float *) x->data, + signs, + nullptr, + nullptr, + dst, + M, + K, + (int) x->ne[1], + (int) x->ne[2], + (int64_t) (x->nb[1] / sizeof(float)), + (int64_t) (x->nb[2] / sizeof(float)), + (int64_t) (x->nb[3] / sizeof(float)) }; + ctx.stream()->parallel_for(sycl::nd_range<2>(sycl::range<2>(M, (K / quant_a_had::HN) * quant_a_had::NT), + sycl::range<2>(1, quant_a_had::NT)), + q); } #else @@ -371,6 +1338,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; } @@ -379,4 +1350,19 @@ void ggml_sycl_pq2_xmx_mul_mat(ggml_backend_sycl_context &, const ggml_tensor *, GGML_ABORT("PQ2_0 XMX path is not built in"); } +ggml_sycl_pq2_xmx_act * ggml_sycl_pq2_xmx_act_quantize(ggml_backend_sycl_context &, const ggml_tensor *, const float *, + ggml_type) { + GGML_ABORT("PQ2_0 XMX path is not built in"); +} + +void ggml_sycl_pq2_xmx_act_free(ggml_sycl_pq2_xmx_act *) {} + +void ggml_sycl_pq2_xmx_mul_mat_act(ggml_backend_sycl_context &, const ggml_tensor *, const ggml_sycl_pq2_xmx_act *, + float *, int, int, const float *) { + GGML_ABORT("PQ2_0 XMX path is not built in"); +} + +void ggml_sycl_pq2_xmx_hadamard_fwht(ggml_backend_sycl_context &, const ggml_tensor *, const float *, float *) { + 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..a13acff85183 100644 --- a/ggml/src/ggml-sycl/pq2_xmx.hpp +++ b/ggml/src/ggml-sycl/pq2_xmx.hpp @@ -2,19 +2,51 @@ #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. - -// ne[0] of a PQ2_0 weight the XMX path accepts: a 2D surface needs a row of at least 64 bytes +// 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 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); // dst = src0 * src1 for a src0 already in the XMX layout; src1 is f32 with contiguous rows, dst is contiguous void ggml_sycl_pq2_xmx_mul_mat(ggml_backend_sycl_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst); + +// Activation quantized for the XMX kernels, reusable by any number of weights of type wtype (PQ2_0 and +// PTQ1_0 read it in different K orders). With had_signs, x is first transformed per row as +// FWHT_1024(x * had_signs), normalized as in ggml_sycl_op_fwht; x->ne[0] must then be a multiple of 1024. +// Freeing returns the memory to the context's xmx_act_pool(), which is safe once the last kernel that reads +// it is submitted. +struct ggml_sycl_pq2_xmx_act; +ggml_sycl_pq2_xmx_act * ggml_sycl_pq2_xmx_act_quantize(ggml_backend_sycl_context & ctx, const ggml_tensor * x, + const float * had_signs, ggml_type wtype); +void ggml_sycl_pq2_xmx_act_free(ggml_sycl_pq2_xmx_act * act); + +// epilogues on the f32 result (TernSYCL postop numbering); other has the layout of dst +enum ggml_sycl_xmx_epi { + GGML_SYCL_XMX_EPI_NONE = 0, + GGML_SYCL_XMX_EPI_SWIGLU = 1, // silu(acc) * other + GGML_SYCL_XMX_EPI_ADD = 2, // acc + other +}; + +// dst = epi(w x act): [w->ne[1], tokens] f32 with row stride ldc, w in the XMX layout. dst and other may alias. +void ggml_sycl_pq2_xmx_mul_mat_act(ggml_backend_sycl_context & ctx, const ggml_tensor * w, + const ggml_sycl_pq2_xmx_act * act, float * dst, int ldc, int epi, + const float * other); + +// dst = FWHT_1024(x * signs) per row, contiguous [x->ne[0], rows] +void ggml_sycl_pq2_xmx_hadamard_fwht(ggml_backend_sycl_context & ctx, const ggml_tensor * x, const float * signs, + float * dst);