Repository navigation
Conversation
Replace the ESIMD kernels of the PQ2_0/PTQ1_0 XMX path with SIMT SYCL kernels ported from Intel's TernSYCL (int2_via_int2_x_int8_dpas, BSD-3-Clause): inline-vISA s2 x s8 DPAS and 2D block I/O, a mat-vec kernel with a K split reduced in SLM for up to 8 tokens, and a GEMM with 256 GRF whose tile depends on the batch size and which splits K over work-groups when a small batch would leave most of the GPU idle. - PQ2_0 layout: uint32 [K/16][N] s2 codes (one 2D block load gives a sub-group 16 columns x 128 weights in the DPAS operand layout), then the fp16 scales [K/128][N]. - PTQ1_0 stays at 1.75 bits: [K/128][7][N] dwords (qs, then qh and the scale). The kernels decode the trits through a 256-entry table in SLM, in byte-major K order so the 10-bit table entries concatenate into the s2 dwords; the activation is quantized in the same permuted order. PTQ1_0 weights no longer expand to 34 bytes a block. For large batches the GEMM decodes each weight block once per work-group and shares it through SLM. - Rows are padded to 16 in the allocation, so any row count works (e.g. 151669-token vocabs). The gate, the layout flag, the capability check and the fallbacks of ggml-org#294 are unchanged.
…othing is shared PTQ1_0 decodes its trits in the kernels, so with few tokens the decode, not the weight reads, set the speed, and the 27B decoded 8 and 16 parallel sequences 3-5% slower than with the expanded weights of ggml-org#294. - The GEMM decodes in registers when a work-group has one sub-group row (WG_M == 1), without the SLM round trip and barrier that only pay off when sub-groups share the decoded block. - PTQ1_0 picks its own tiles: 16x16 and 32x16 per sub-group for 5-32 tokens (each decoded block feeds 2-4 row blocks), 32x32 with one sub-group row for 33-64 tokens with K >= 6144, and the 4-row mat-vec for 2 tokens (the 2-row variant is slower). Arc Pro B50, Bonsai 2 27B PTQ1_0, llama-batched-bench decode t/s at 8 / 16 / 32 sequences: 77.4 / 93.1 / 110.7 -> 82.2 / 105.7 / 119.7 (prism: 80.1 / 98.1 / 110.8).
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Overview
Replaces the ESIMD kernels of the PQ2_0/PTQ1_0 XMX path from #294 with SIMT SYCL kernels ported from Intel's TernSYCL (
int2_via_int2_x_int8_dpas, BSD-3-Clause, notice kept inpq2_xmx.cpp), and keeps PTQ1_0 at 1.75 bits in VRAM instead of expanding it to the PQ2_0 layout.The gate, the
xmx_pq2layout flag, the 16-wide-DPAS capability check, the AOT guard in CMake, theget_rowsasserts and every fallback of #294 are unchanged. The diff ispq2_xmx.cpp/.hppplus the allocation-size rule inggml-sycl.cpp.How it works:
[K/16][N]s2 codes, then the fp16 scales[K/128][N]. One 2D block load gives a sub-group 16 columns x 128 weights directly in the DPAS operand layout.[K/128][7][N]dwords (the 24 qs bytes, then qh and the scale), the same size as the file. The kernels decode the base-3 trits through a 256-entry table in SLM. They decode in byte-major K order (the 5 trits of a byte next to each other), so the 10-bit table entries of consecutive bytes concatenate straight into the s2 dwords; the activation quantizer writes the int8 activations in the same permuted order, so the dot products are unchanged. For large batches each weight block is decoded once per work-group and shared through SLM.Two commits: the kernels, and a follow-up that picks PTQ1_0's own tiles for small batches (2-64 tokens) and decodes in registers when no sub-groups share a decoded block. Without it, the in-kernel trit decode made PTQ1_0 3-5% slower than #294 at 8-16 parallel sequences.
Results
Intel Arc Pro B50 (16 GB), oneAPI 2026.1, Level Zero 1.14.37020, Linux 7.0
xedriver,-ngl 99, Ternary Bonsai models. "prism" is the currentprismhead (6bfcd79, #294 path). All columns measured in the same session.llama-bench, t/s:-fa 1)-fa 1)(pp128 of the 1.7B swings by +-1700 t/s between repetitions on both builds, so it is left out.)
Mat-mul alone, weights rotated through 8 copies so every call reads DRAM, including the activation quantization, us per call:
Parallel decode (
llama-batched-bench -npp 128 -ntg 64), aggregate decode t/s, prism / this PR:The 27B's gated delta net layers spend more per sequence in concat and get_rows than in the mat-muls at higher counts; with #326 on top (measured on the first commit) 32 sequences reach 151.9 (PQ2_0) and 140.7 (PTQ1_0) t/s.
PTQ1_0 weights stay at 1.75 bits, where #294 reserves 34 bytes per 28-byte block (about +1.2 GB for the 27B). The mat-muls on the small 1.7B shape (2048 x 6144) are still 5-38% slower than #294 for PTQ1_0 at 2-16 tokens.
Speculative decoding with the multi-token-prediction head (sudoingx/Ternary-Bonsai-2-27B-PTQ1_0-MTP-GGUF, the Qwen MTP head grafted onto the PTQ1_0 file),
llama-server, 3 chat prompts (prose, code, bash) x 256 tokens, greedy, thinking off, decode t/s:--spec-type draft-mtp --spec-draft-n-max 1--spec-type draft-mtp --spec-draft-n-max 2Correctness
test-backend-ops -b SYCL0, run per op over all 130 ops: MUL_MAT 1401/1401, MUL_MAT_ID 1035/1035, MUL_MAT_VEC_FUSION 506/506, GET_ROWS 119/119. The ops that do not pass fully (CONV_2D 2010/2026, FLASH_ATTN_EXT 4070/4073, LIGHTNING_INDEXER 117-120/156 varying between runs, ROLL 1/2; CPY and SET_ROWS abort) do the same on prism.-ub6 / 16 / 32 / 64 / 128 / 512 gives 6.2386 / 6.2358 / 6.2394 / 6.2375 / 6.2420 / 6.2347 (+-0.42).GGML_SYCL_NO_PQ2_XMXdefined (the AOT guard for devices without 16-wide DPAS),pq2_xmx.cppcompiles to the stubs.-c 512, 5 chunks of the repository'sdocs/*.md), prism -> this PR:(The two 27B files hold the same ternary weights, so they give the same numbers.)
-dev none --no-op-offload), prism vs this PR: b1 0.00127 vs 0.00130, b512 0.00127 vs 0.00120; same top token 98.4% vs 98.0% (b1), 97.6% vs 97.8% (b512). Within the error bars.Additional information
pq2_xmx.cppare compiled together, so a kernel the backend compiler rejects breaks the others too; every tile used here is exercised by the MUL_MAT cases above.Requirements