whispercpp 1.3.6 → 1.3.8
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- checksums.yaml +4 -4
- data/.document +3 -0
- data/.rdoc_options +2 -0
- data/README.md +43 -9
- data/Rakefile +18 -3
- data/ext/dependencies.rb +10 -4
- data/ext/dependencies_for_windows.rb +17 -0
- data/ext/extconf.rb +20 -8
- data/ext/options.rb +54 -14
- data/ext/options_for_windows.rb +51 -0
- data/ext/ruby_whisper.c +35 -42
- data/ext/ruby_whisper.h +141 -0
- data/ext/ruby_whisper_context.c +157 -29
- data/ext/ruby_whisper_log_queue.c +180 -0
- data/ext/ruby_whisper_log_settable.h +46 -0
- data/ext/ruby_whisper_parakeet.c +49 -0
- data/ext/ruby_whisper_parakeet_context.c +304 -0
- data/ext/ruby_whisper_parakeet_context_params.c +117 -0
- data/ext/ruby_whisper_parakeet_model.c +84 -0
- data/ext/ruby_whisper_parakeet_params.c +548 -0
- data/ext/ruby_whisper_parakeet_segment.c +157 -0
- data/ext/ruby_whisper_parakeet_token.c +188 -0
- data/ext/ruby_whisper_parakeet_transcribe.cpp +58 -0
- data/ext/ruby_whisper_params.c +265 -73
- data/ext/ruby_whisper_segment.c +6 -6
- data/ext/ruby_whisper_transcribe.cpp +23 -15
- data/ext/ruby_whisper_vad_context.c +30 -10
- data/ext/ruby_whisper_vad_context_detect.cpp +8 -9
- data/ext/ruby_whisper_vad_params.c +4 -4
- data/ext/ruby_whisper_vad_segment.c +2 -2
- data/ext/sources/CMakeLists.txt +42 -3
- data/ext/sources/CMakePresets.json +95 -0
- data/ext/sources/cmake/parakeet-config.cmake.in +30 -0
- data/ext/sources/cmake/parakeet.pc.in +10 -0
- data/ext/sources/cmake/whisper.pc.in +2 -2
- data/ext/sources/examples/CMakeLists.txt +4 -2
- data/ext/sources/examples/bench/bench.cpp +1 -1
- data/ext/sources/examples/cli/cli.cpp +52 -10
- data/ext/sources/examples/common-ggml.cpp +4 -0
- data/ext/sources/examples/common-whisper.cpp +139 -67
- data/ext/sources/examples/common-whisper.h +11 -0
- data/ext/sources/examples/ffmpeg-transcode.cpp +211 -341
- data/ext/sources/examples/parakeet-cli/CMakeLists.txt +8 -0
- data/ext/sources/examples/parakeet-cli/parakeet-cli.cpp +243 -0
- data/ext/sources/examples/parakeet-quantize/CMakeLists.txt +7 -0
- data/ext/sources/examples/parakeet-quantize/parakeet-quantize.cpp +230 -0
- data/ext/sources/examples/server/server.cpp +199 -163
- data/ext/sources/examples/vad-speech-segments/speech.cpp +3 -2
- data/ext/sources/ggml/CMakeLists.txt +21 -14
- data/ext/sources/ggml/cmake/FindNCCL.cmake +36 -0
- data/ext/sources/ggml/cmake/ggml-config.cmake.in +12 -2
- data/ext/sources/ggml/include/ggml-alloc.h +1 -0
- data/ext/sources/ggml/include/ggml-backend.h +72 -10
- data/ext/sources/ggml/include/ggml-cuda.h +2 -2
- data/ext/sources/ggml/include/ggml-rpc.h +3 -3
- data/ext/sources/ggml/include/ggml-sycl.h +8 -0
- data/ext/sources/ggml/include/ggml.h +103 -9
- data/ext/sources/ggml/include/gguf.h +10 -2
- data/ext/sources/ggml/src/CMakeLists.txt +30 -6
- data/ext/sources/ggml/src/ggml-alloc.c +5 -1
- data/ext/sources/ggml/src/ggml-backend-impl.h +22 -2
- data/ext/sources/ggml/src/ggml-backend-meta.cpp +2266 -0
- data/ext/sources/ggml/src/ggml-backend-reg.cpp +12 -0
- data/ext/sources/ggml/src/ggml-backend.cpp +110 -9
- data/ext/sources/ggml/src/ggml-blas/ggml-blas.cpp +4 -0
- data/ext/sources/ggml/src/ggml-cann/aclnn_ops.cpp +672 -257
- data/ext/sources/ggml/src/ggml-cann/aclnn_ops.h +71 -0
- data/ext/sources/ggml/src/ggml-cann/common.h +20 -10
- data/ext/sources/ggml/src/ggml-cann/ggml-cann.cpp +211 -30
- data/ext/sources/ggml/src/ggml-common.h +24 -2
- data/ext/sources/ggml/src/ggml-cpu/CMakeLists.txt +59 -30
- data/ext/sources/ggml/src/ggml-cpu/amx/amx.cpp +2 -0
- data/ext/sources/ggml/src/ggml-cpu/amx/mmq.cpp +21 -22
- data/ext/sources/ggml/src/ggml-cpu/arch/arm/quants.c +194 -11
- data/ext/sources/ggml/src/ggml-cpu/arch/arm/repack.cpp +65 -0
- data/ext/sources/ggml/src/ggml-cpu/arch/loongarch/quants.c +151 -1
- data/ext/sources/ggml/src/ggml-cpu/arch/powerpc/quants.c +0 -1
- data/ext/sources/ggml/src/ggml-cpu/arch/riscv/quants.c +4279 -1292
- data/ext/sources/ggml/src/ggml-cpu/arch/riscv/repack.cpp +5 -35
- data/ext/sources/ggml/src/ggml-cpu/arch/s390/quants.c +0 -1
- data/ext/sources/ggml/src/ggml-cpu/arch/wasm/quants.c +72 -1
- data/ext/sources/ggml/src/ggml-cpu/arch/x86/quants.c +319 -31
- data/ext/sources/ggml/src/ggml-cpu/arch/x86/repack.cpp +1 -1
- data/ext/sources/ggml/src/ggml-cpu/arch-fallback.h +12 -2
- data/ext/sources/ggml/src/ggml-cpu/cmake/FindSMTIME.cmake +32 -0
- data/ext/sources/ggml/src/ggml-cpu/ggml-cpu-impl.h +10 -0
- data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.c +109 -5
- data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.cpp +2 -0
- data/ext/sources/ggml/src/ggml-cpu/kleidiai/kleidiai.cpp +146 -134
- data/ext/sources/ggml/src/ggml-cpu/llamafile/sgemm.cpp +107 -82
- data/ext/sources/ggml/src/ggml-cpu/ops.cpp +501 -119
- data/ext/sources/ggml/src/ggml-cpu/ops.h +3 -0
- data/ext/sources/ggml/src/ggml-cpu/quants.c +106 -0
- data/ext/sources/ggml/src/ggml-cpu/quants.h +6 -0
- data/ext/sources/ggml/src/ggml-cpu/repack.cpp +3 -0
- data/ext/sources/ggml/src/ggml-cpu/simd-gemm.h +91 -1
- data/ext/sources/ggml/src/ggml-cpu/simd-mappings.h +14 -16
- data/ext/sources/ggml/src/ggml-cpu/spacemit/ime.cpp +1402 -687
- data/ext/sources/ggml/src/ggml-cpu/spacemit/ime.h +8 -0
- data/ext/sources/ggml/src/ggml-cpu/spacemit/ime1_kernels.cpp +597 -2766
- data/ext/sources/ggml/src/ggml-cpu/spacemit/ime2_kernels.cpp +5768 -0
- data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_env.cpp +320 -0
- data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_env.h +55 -0
- data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_kernels.h +182 -19
- data/ext/sources/ggml/src/ggml-cpu/spacemit/repack.cpp +1795 -0
- data/ext/sources/ggml/src/ggml-cpu/spacemit/repack.h +14 -0
- data/ext/sources/ggml/src/ggml-cpu/spacemit/rvv_kernels.cpp +3178 -0
- data/ext/sources/ggml/src/ggml-cpu/spacemit/rvv_kernels.h +95 -0
- data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_barrier.h +34 -0
- data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_mem_pool.cpp +760 -0
- data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_mem_pool.h +32 -0
- data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_tcm.h +409 -0
- data/ext/sources/ggml/src/ggml-cpu/vec.cpp +39 -55
- data/ext/sources/ggml/src/ggml-cpu/vec.h +225 -240
- data/ext/sources/ggml/src/ggml-cuda/CMakeLists.txt +17 -7
- data/ext/sources/ggml/src/ggml-cuda/allreduce.cu +971 -0
- data/ext/sources/ggml/src/ggml-cuda/allreduce.cuh +29 -0
- data/ext/sources/ggml/src/ggml-cuda/argsort.cu +62 -26
- data/ext/sources/ggml/src/ggml-cuda/binbcast.cu +134 -64
- data/ext/sources/ggml/src/ggml-cuda/binbcast.cuh +1 -0
- data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cu +81 -0
- data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cuh +3 -0
- data/ext/sources/ggml/src/ggml-cuda/common.cuh +246 -28
- data/ext/sources/ggml/src/ggml-cuda/concat.cu +134 -116
- data/ext/sources/ggml/src/ggml-cuda/conv-transpose-1d.cu +14 -12
- data/ext/sources/ggml/src/ggml-cuda/conv2d-transpose.cu +45 -21
- data/ext/sources/ggml/src/ggml-cuda/conv2d-transpose.cuh +1 -0
- data/ext/sources/ggml/src/ggml-cuda/convert.cu +139 -34
- data/ext/sources/ggml/src/ggml-cuda/convert.cuh +10 -0
- data/ext/sources/ggml/src/ggml-cuda/cpy.cu +88 -29
- data/ext/sources/ggml/src/ggml-cuda/dequantize.cuh +22 -0
- data/ext/sources/ggml/src/ggml-cuda/fattn-common.cuh +287 -49
- data/ext/sources/ggml/src/ggml-cuda/fattn-mma-f16.cuh +335 -130
- data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cu +12 -0
- data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cuh +127 -24
- data/ext/sources/ggml/src/ggml-cuda/fattn-vec.cuh +40 -15
- data/ext/sources/ggml/src/ggml-cuda/fattn-wmma-f16.cu +18 -9
- data/ext/sources/ggml/src/ggml-cuda/fattn.cu +169 -60
- data/ext/sources/ggml/src/ggml-cuda/fattn.cuh +2 -0
- data/ext/sources/ggml/src/ggml-cuda/fwht.cu +101 -0
- data/ext/sources/ggml/src/ggml-cuda/fwht.cuh +4 -0
- data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cu +109 -45
- data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cuh +10 -0
- data/ext/sources/ggml/src/ggml-cuda/getrows.cu +48 -23
- data/ext/sources/ggml/src/ggml-cuda/ggml-cuda.cu +2034 -2104
- data/ext/sources/ggml/src/ggml-cuda/im2col.cu +32 -29
- data/ext/sources/ggml/src/ggml-cuda/mean.cu +4 -2
- data/ext/sources/ggml/src/ggml-cuda/mma.cuh +242 -195
- data/ext/sources/ggml/src/ggml-cuda/mmf.cuh +3 -3
- data/ext/sources/ggml/src/ggml-cuda/mmq.cu +25 -12
- data/ext/sources/ggml/src/ggml-cuda/mmq.cuh +502 -423
- data/ext/sources/ggml/src/ggml-cuda/mmvf.cu +19 -12
- data/ext/sources/ggml/src/ggml-cuda/mmvq.cu +562 -97
- data/ext/sources/ggml/src/ggml-cuda/mmvq.cuh +6 -1
- data/ext/sources/ggml/src/ggml-cuda/norm.cu +36 -10
- data/ext/sources/ggml/src/ggml-cuda/out-prod.cu +66 -7
- data/ext/sources/ggml/src/ggml-cuda/quantize.cu +133 -26
- data/ext/sources/ggml/src/ggml-cuda/quantize.cuh +1 -1
- data/ext/sources/ggml/src/ggml-cuda/reduce_rows.cuh +5 -1
- data/ext/sources/ggml/src/ggml-cuda/rope.cu +11 -4
- data/ext/sources/ggml/src/ggml-cuda/scale.cu +4 -1
- data/ext/sources/ggml/src/ggml-cuda/set-rows.cu +78 -10
- data/ext/sources/ggml/src/ggml-cuda/snake.cu +72 -0
- data/ext/sources/ggml/src/ggml-cuda/snake.cuh +8 -0
- data/ext/sources/ggml/src/ggml-cuda/softcap.cu +4 -1
- data/ext/sources/ggml/src/ggml-cuda/ssm-conv.cu +45 -13
- data/ext/sources/ggml/src/ggml-cuda/ssm-conv.cuh +1 -1
- data/ext/sources/ggml/src/ggml-cuda/ssm-scan.cu +40 -18
- data/ext/sources/ggml/src/ggml-cuda/sumrows.cu +8 -4
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_16.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_32.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_8.cu +2 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_2.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_4.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_16.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_32.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_4.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_8.cu +2 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_32-ncols2_2.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_16.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_2.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_4.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_8.cu +2 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_2.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_4.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_8.cu +2 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq192-dv128.cu +5 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq320-dv256.cu +5 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq512-dv512.cu +5 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-bf16.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-f16.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q4_0.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q4_1.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q5_0.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q5_1.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q8_0.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-f16-bf16.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q4_0-bf16.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q4_1-bf16.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q5_0-bf16.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q5_1-bf16.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q8_0-bf16.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/mmq-instance-nvfp4.cu +5 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/mmq-instance-q1_0.cu +5 -0
- data/ext/sources/ggml/src/ggml-cuda/top-k.cu +5 -4
- data/ext/sources/ggml/src/ggml-cuda/topk-moe.cu +33 -24
- data/ext/sources/ggml/src/ggml-cuda/unary.cu +31 -2
- data/ext/sources/ggml/src/ggml-cuda/unary.cuh +2 -0
- data/ext/sources/ggml/src/ggml-cuda/vecdotq.cuh +80 -0
- data/ext/sources/ggml/src/ggml-cuda/vendors/cuda.h +7 -2
- data/ext/sources/ggml/src/ggml-cuda/vendors/hip.h +23 -4
- data/ext/sources/ggml/src/ggml-cuda/vendors/musa.h +4 -0
- data/ext/sources/ggml/src/ggml-hexagon/CMakeLists.txt +1 -5
- data/ext/sources/ggml/src/ggml-hexagon/ggml-hexagon.cpp +2788 -1762
- data/ext/sources/ggml/src/ggml-hexagon/htp/CMakeLists.txt +13 -4
- data/ext/sources/ggml/src/ggml-hexagon/htp/act-ops.c +53 -84
- data/ext/sources/ggml/src/ggml-hexagon/htp/argsort-ops.c +25 -12
- data/ext/sources/ggml/src/ggml-hexagon/htp/binary-ops.c +165 -184
- data/ext/sources/ggml/src/ggml-hexagon/htp/cmake-toolchain.cmake +17 -19
- data/ext/sources/ggml/src/ggml-hexagon/htp/concat-ops.c +277 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/cpy-ops.c +170 -127
- data/ext/sources/ggml/src/ggml-hexagon/htp/cumsum-ops.c +270 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/diag-ops.c +216 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/fill-ops.c +123 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +1774 -396
- data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.h +303 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c +1148 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/get-rows-ops.c +148 -42
- data/ext/sources/ggml/src/ggml-hexagon/htp/hex-common.h +80 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.c +2 -2
- data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.h +255 -62
- data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dump.h +9 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hex-profile.h +64 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hex-utils.h +25 -21
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h +555 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h +1303 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.c +167 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.h +157 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-utils.h +222 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ctx.h +104 -13
- data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ops.h +222 -57
- data/ext/sources/ggml/src/ggml-hexagon/htp/htp-vtcm.h +19 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/htp_iface.idl +10 -3
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-base.h +78 -26
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-copy.h +27 -10
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-div.h +63 -23
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-exp.h +48 -8
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h +232 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-flash-attn.h +47 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-log.h +65 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-flat.h +1511 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h +1200 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-pow.h +42 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-repl.h +74 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h +40 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sin-cos.h +90 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-utils.h +5 -8
- data/ext/sources/ggml/src/ggml-hexagon/htp/main.c +625 -816
- data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.c +3052 -2166
- data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.h +650 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/pad-ops.c +547 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/repeat-ops.c +148 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/rope-ops.c +337 -106
- data/ext/sources/ggml/src/ggml-hexagon/htp/set-rows-ops.c +59 -37
- data/ext/sources/ggml/src/ggml-hexagon/htp/softmax-ops.c +121 -133
- data/ext/sources/ggml/src/ggml-hexagon/htp/solve-tri-ops.c +267 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/ssm-conv.c +245 -151
- data/ext/sources/ggml/src/ggml-hexagon/htp/sum-rows-ops.c +6 -6
- data/ext/sources/ggml/src/ggml-hexagon/htp/unary-ops.c +719 -45
- data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.c +15 -3
- data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.h +8 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp-opnode.h +390 -0
- data/ext/sources/ggml/src/ggml-hexagon/libggml-htp.inf +3 -5
- data/ext/sources/ggml/src/ggml-hip/CMakeLists.txt +27 -9
- data/ext/sources/ggml/src/ggml-impl.h +6 -1
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.cpp +207 -18
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.h +36 -2
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.m +186 -29
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-impl.h +118 -0
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.cpp +322 -21
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.h +4 -0
- data/ext/sources/ggml/src/ggml-metal/ggml-metal.cpp +39 -26
- data/ext/sources/ggml/src/ggml-metal/ggml-metal.metal +1226 -467
- data/ext/sources/ggml/src/ggml-musa/CMakeLists.txt +5 -6
- data/ext/sources/ggml/src/ggml-opencl/CMakeLists.txt +67 -5
- data/ext/sources/ggml/src/ggml-opencl/fa_tune.h +92 -0
- data/ext/sources/ggml/src/ggml-opencl/ggml-opencl.cpp +16290 -6246
- data/ext/sources/ggml/src/ggml-opencl/kernels/concat.cl +67 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/cpy.cl +59 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/cvt.cl +1997 -92
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f16.cl +81 -41
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32.cl +88 -39
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl +1995 -96
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q4_0.cl +1615 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q8_0.cl +1486 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_pre_f16.cl +156 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gated_delta_net.cl +249 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_mxfp4_f32_ns.cl +374 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_0_f32_ns.cl +324 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_1_f32_ns.cl +326 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_k_f32_ns.cl +348 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_0_f32_ns.cl +328 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_1_f32_ns.cl +330 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_k_f32_ns.cl +356 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q6_k_f32_ns.cl +335 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_iq4_nl_f32.cl +150 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q1_0_f32.cl +94 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/{mul_mat_Ab_Bi_8x4.cl → gemm_noshuffle_q4_0_f32.cl} +1 -1
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q4_k_f32.cl +172 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_0_f32.cl +131 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_1_f32.cl +134 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_k_f32.cl +176 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q6_k_f32.cl +140 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/{mul_mm_q8_0_f32_8x4.cl → gemm_noshuffle_q8_0_f32.cl} +1 -1
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_xmem_f16_f32_os8.cl +233 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_mxfp4_f32_ns.cl +165 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_0_f32_ns.cl +120 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_1_f32_ns.cl +123 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_k_f32_ns.cl +155 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_0_f32_ns.cl +123 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_1_f32_ns.cl +125 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_k_f32_ns.cl +160 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q6_k_f32_ns.cl +141 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_iq4_nl_f32.cl +302 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q1_0_f32.cl +121 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle_general.cl → gemv_noshuffle_q4_0_f32.cl} +5 -5
- data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle.cl → gemv_noshuffle_q4_0_f32_spec.cl} +5 -5
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32.cl +318 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_0_f32.cl +291 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_1_f32.cl +294 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_k_f32.cl +326 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32.cl +293 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle_general_q8_0_f32.cl → gemv_noshuffle_q8_0_f32.cl} +1 -1
- data/ext/sources/ggml/src/ggml-opencl/kernels/get_rows.cl +15 -9
- data/ext/sources/ggml/src/ggml-opencl/kernels/moe_reorder_b.cl +30 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/moe_sort_by_expert.cl +82 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_iq4_nl_f32_l4_lm.cl +171 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q1_0_f32_l4_lm.cl +156 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q4_k_f32_l4_lm.cl +179 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_0_f32_l4_lm.cl +173 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_1_f32_l4_lm.cl +175 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_k_f32_l4_lm.cl +192 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_l4.cl +1149 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_iq4_nl_f32.cl +164 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_iq4_nl_f32_flat.cl +202 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32.cl +141 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32_flat.cl +190 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q4_k_f32_flat.cl +196 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_0_f32.cl +241 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_0_f32_flat.cl +243 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_1_f32.cl +243 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_1_f32_flat.cl +247 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32.cl +187 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32_flat.cl +203 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q6_k_f32_flat.cl +48 -64
- data/ext/sources/ggml/src/ggml-opencl/kernels/norm.cl +5 -2
- data/ext/sources/ggml/src/ggml-opencl/kernels/set_rows.cl +500 -0
- data/ext/sources/ggml/src/ggml-opencl/libdl.h +79 -0
- data/ext/sources/ggml/src/ggml-openvino/.clang-format +0 -5
- data/ext/sources/ggml/src/ggml-openvino/CMakeLists.txt +2 -4
- data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.cpp +740 -127
- data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.h +76 -23
- data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.cpp +75 -14
- data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.h +29 -8
- data/ext/sources/ggml/src/ggml-openvino/ggml-openvino.cpp +339 -69
- data/ext/sources/ggml/src/ggml-openvino/ggml-quants.cpp +330 -192
- data/ext/sources/ggml/src/ggml-openvino/ggml-quants.h +10 -4
- data/ext/sources/ggml/src/ggml-openvino/openvino/decoder.h +56 -16
- data/ext/sources/ggml/src/ggml-openvino/openvino/frontend.h +1 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/input_model.h +4 -4
- data/ext/sources/ggml/src/ggml-openvino/openvino/node_context.h +94 -37
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/add_id.cpp +76 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/argsort.cpp +47 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/clamp.cpp +33 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/concat.cpp +48 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/cont.cpp +8 -16
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/cpy.cpp +14 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/div.cpp +146 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp +108 -21
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp +282 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.hpp +65 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/get_rows.cpp +2 -9
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_geglu.cpp +21 -7
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp +41 -8
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/im2col.cpp +120 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/l2_norm.cpp +44 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/mul_mat_id.cpp +226 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/mulmat.cpp +19 -9
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/norm.cpp +58 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/pad.cpp +95 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/permute.cpp +58 -13
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/repeat.cpp +74 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/reshape.cpp +13 -6
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/rms_norm.cpp +1 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/rope.cpp +161 -39
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/set_rows.cpp +3 -3
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/softmax.cpp +126 -49
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/ssm_conv.cpp +59 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/sum_rows.cpp +27 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/transpose.cpp +32 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_silu.cpp +1 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_softplus.cpp +38 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/view.cpp +90 -25
- data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.cpp +41 -22
- data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.h +18 -4
- data/ext/sources/ggml/src/ggml-openvino/openvino/pass/mark_decompression_convert_constant_folding.h +1 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/rt_info/weightless_caching_attributes.hpp +41 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.cpp +70 -43
- data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.h +5 -4
- data/ext/sources/ggml/src/ggml-openvino/openvino/utils.cpp +612 -36
- data/ext/sources/ggml/src/ggml-openvino/openvino/utils.h +29 -26
- data/ext/sources/ggml/src/ggml-openvino/utils.cpp +460 -114
- data/ext/sources/ggml/src/ggml-openvino/utils.h +32 -9
- data/ext/sources/ggml/src/ggml-opt.cpp +1 -0
- data/ext/sources/ggml/src/ggml-quants.c +365 -114
- data/ext/sources/ggml/src/ggml-quants.h +6 -0
- data/ext/sources/ggml/src/ggml-rpc/CMakeLists.txt +24 -0
- data/ext/sources/ggml/src/ggml-rpc/ggml-rpc.cpp +167 -311
- data/ext/sources/ggml/src/ggml-rpc/transport.cpp +683 -0
- data/ext/sources/ggml/src/ggml-rpc/transport.h +34 -0
- data/ext/sources/ggml/src/ggml-sycl/CMakeLists.txt +50 -4
- data/ext/sources/ggml/src/ggml-sycl/add-id.cpp +1 -1
- data/ext/sources/ggml/src/ggml-sycl/backend.hpp +5 -1
- data/ext/sources/ggml/src/ggml-sycl/binbcast.cpp +12 -0
- data/ext/sources/ggml/src/ggml-sycl/col2im-1d.cpp +102 -0
- data/ext/sources/ggml/src/ggml-sycl/col2im-1d.hpp +8 -0
- data/ext/sources/ggml/src/ggml-sycl/common.cpp +72 -2
- data/ext/sources/ggml/src/ggml-sycl/common.hpp +59 -2
- data/ext/sources/ggml/src/ggml-sycl/concat.cpp +21 -1
- data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.cpp +158 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.hpp +10 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.cpp +125 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.hpp +10 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d.cpp +150 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d.hpp +10 -0
- data/ext/sources/ggml/src/ggml-sycl/conv3d.cpp +224 -0
- data/ext/sources/ggml/src/ggml-sycl/conv3d.hpp +8 -0
- data/ext/sources/ggml/src/ggml-sycl/convert.cpp +121 -13
- data/ext/sources/ggml/src/ggml-sycl/convert.hpp +9 -0
- data/ext/sources/ggml/src/ggml-sycl/cpy.cpp +706 -0
- data/ext/sources/ggml/src/ggml-sycl/cpy.hpp +281 -0
- data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.cpp +255 -0
- data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.hpp +7 -0
- data/ext/sources/ggml/src/ggml-sycl/cumsum.cpp +148 -0
- data/ext/sources/ggml/src/ggml-sycl/cumsum.hpp +5 -0
- data/ext/sources/ggml/src/ggml-sycl/dequantize.hpp +678 -0
- data/ext/sources/ggml/src/ggml-sycl/diag.cpp +67 -0
- data/ext/sources/ggml/src/ggml-sycl/diag.hpp +5 -0
- data/ext/sources/ggml/src/ggml-sycl/dmmv.cpp +997 -244
- data/ext/sources/ggml/src/ggml-sycl/dpct/helper.hpp +15 -7
- data/ext/sources/ggml/src/ggml-sycl/element_wise.cpp +215 -204
- data/ext/sources/ggml/src/ggml-sycl/element_wise.hpp +2 -2
- data/ext/sources/ggml/src/ggml-sycl/fattn-buffers.cpp +56 -0
- data/ext/sources/ggml/src/ggml-sycl/fattn-buffers.hpp +63 -0
- data/ext/sources/ggml/src/ggml-sycl/fattn-common.hpp +7 -5
- data/ext/sources/ggml/src/ggml-sycl/fattn-tile.cpp +4 -0
- data/ext/sources/ggml/src/ggml-sycl/fattn-tile.hpp +76 -168
- data/ext/sources/ggml/src/ggml-sycl/fattn-vec.hpp +7 -0
- data/ext/sources/ggml/src/ggml-sycl/fattn.cpp +3 -1
- data/ext/sources/ggml/src/ggml-sycl/fill.cpp +55 -0
- data/ext/sources/ggml/src/ggml-sycl/fill.hpp +5 -0
- data/ext/sources/ggml/src/ggml-sycl/gated_delta_net.cpp +69 -31
- data/ext/sources/ggml/src/ggml-sycl/gated_delta_net.hpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/gemm.hpp +3 -0
- data/ext/sources/ggml/src/ggml-sycl/getrows.cpp +79 -3
- data/ext/sources/ggml/src/ggml-sycl/ggml-sycl.cpp +1758 -455
- data/ext/sources/ggml/src/ggml-sycl/im2col.cpp +353 -89
- data/ext/sources/ggml/src/ggml-sycl/im2col.hpp +5 -3
- data/ext/sources/ggml/src/ggml-sycl/mmvq.cpp +1542 -39
- data/ext/sources/ggml/src/ggml-sycl/mmvq.hpp +33 -0
- data/ext/sources/ggml/src/ggml-sycl/norm.cpp +103 -49
- data/ext/sources/ggml/src/ggml-sycl/outprod.cpp +45 -9
- data/ext/sources/ggml/src/ggml-sycl/pad.cpp +27 -27
- data/ext/sources/ggml/src/ggml-sycl/pool.cpp +185 -0
- data/ext/sources/ggml/src/ggml-sycl/pool.hpp +22 -0
- data/ext/sources/ggml/src/ggml-sycl/presets.hpp +3 -1
- data/ext/sources/ggml/src/ggml-sycl/quants.hpp +71 -0
- data/ext/sources/ggml/src/ggml-sycl/set_rows.cpp +17 -3
- data/ext/sources/ggml/src/ggml-sycl/softmax.cpp +9 -10
- data/ext/sources/ggml/src/ggml-sycl/solve_tri.cpp +172 -0
- data/ext/sources/ggml/src/ggml-sycl/solve_tri.hpp +8 -0
- data/ext/sources/ggml/src/ggml-sycl/ssm_conv.cpp +6 -1
- data/ext/sources/ggml/src/ggml-sycl/ssm_scan.cpp +156 -0
- data/ext/sources/ggml/src/ggml-sycl/ssm_scan.hpp +5 -0
- data/ext/sources/ggml/src/ggml-sycl/sycl_hw.cpp +62 -10
- data/ext/sources/ggml/src/ggml-sycl/sycl_hw.hpp +18 -6
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-tile-instance-dkq512-dv512.cpp +6 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-f16.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q4_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q4_1.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q5_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q5_1.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q8_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-f16.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q4_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q4_1.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q5_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q5_1.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q8_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-f16.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q4_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q4_1.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q5_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q5_1.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q8_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-f16.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q4_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q4_1.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q5_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q5_1.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q8_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-f16.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q4_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q4_1.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q5_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q5_1.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q8_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-f16.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q4_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q4_1.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q5_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q5_1.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q8_0.cpp +1 -0
- data/ext/sources/ggml/src/ggml-sycl/type.hpp +112 -0
- data/ext/sources/ggml/src/ggml-sycl/upscale.cpp +410 -0
- data/ext/sources/ggml/src/ggml-sycl/upscale.hpp +9 -0
- data/ext/sources/ggml/src/ggml-sycl/vecdotq.hpp +242 -45
- data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend-buffer.cpp +4 -0
- data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend-device.cpp +2 -0
- data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend.cpp +2 -0
- data/ext/sources/ggml/src/ggml-virtgpu/virtgpu-shm.cpp +1 -0
- data/ext/sources/ggml/src/ggml-virtgpu/virtgpu.cpp +1 -0
- data/ext/sources/ggml/src/ggml-virtgpu/virtgpu.h +0 -2
- data/ext/sources/ggml/src/ggml-vulkan/CMakeLists.txt +16 -0
- data/ext/sources/ggml/src/ggml-vulkan/ggml-vulkan.cpp +2843 -700
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/CMakeLists.txt +4 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/col2im_1d.comp +61 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/contig_copy.comp +6 -2
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp +146 -13
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp +431 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy.comp +3 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy_from_quant.comp +1 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy_to_quant.comp +25 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl +88 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl +643 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_nvfp4.comp +32 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q1_0.comp +29 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/diag.comp +3 -4
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dot_product_funcs.glsl +27 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/feature-tests/coopmat2_decode_vector.comp +7 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +198 -48
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl +60 -59
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +116 -113
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp +122 -31
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_dequant.glsl +131 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_mmq_funcs.glsl +203 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/fwht.comp +115 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gated_delta_net.comp +125 -64
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_binary_head.glsl +0 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl +21 -19
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/get_rows_back.comp +25 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl +29 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_main.glsl +17 -11
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/im2col.comp +76 -54
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/im2col_3d.comp +0 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/l2_norm.comp +4 -7
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/log.comp +0 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec.comp +122 -27
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iface.glsl +6 -6
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q2_k.comp +1 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q4_k.comp +1 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q5_k.comp +1 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp +22 -24
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_funcs.glsl +88 -55
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +42 -40
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +49 -15
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +222 -171
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_funcs.glsl +8 -8
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_shmem_types.glsl +24 -9
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/multi_add.comp +0 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/norm.comp +10 -10
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/repeat_back.comp +3 -3
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/roll.comp +3 -3
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_funcs.glsl +5 -2
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_head.glsl +0 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_params.glsl +3 -2
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/snake.comp +49 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ssm_conv.comp +11 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tri.comp +3 -4
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl +79 -2
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/unary.comp +168 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +282 -211
- data/ext/sources/ggml/src/ggml-webgpu/CMakeLists.txt +5 -2
- data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +2209 -283
- data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu.cpp +2618 -1416
- data/ext/sources/ggml/src/ggml-webgpu/pre_wgsl.hpp +37 -7
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/add_id.wgsl +64 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/binary.wgsl +8 -7
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl +90 -95
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl +19 -1
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/conv2d.wgsl +165 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{cpy.tmpl.wgsl → cpy.wgsl} +25 -50
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl +107 -184
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_quant_staging.tmpl +124 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_tile.wgsl +397 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_blk.wgsl +101 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_reduce.wgsl +84 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl +619 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/gated_delta_net.wgsl +149 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl +204 -78
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/glu.wgsl +155 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/im2col.wgsl +101 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl +805 -526
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id.wgsl +195 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_gather.wgsl +52 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_vec.wgsl +154 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_reg_tile.wgsl +8 -6
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_subgroup_matrix.wgsl +5 -1
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl +90 -413
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl +1553 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl +297 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quant_inner_loops.tmpl +21 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl +178 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm_mul.wgsl +152 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{rope.tmpl.wgsl → rope.wgsl} +71 -142
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/row_norm.wgsl +153 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/scale.wgsl +6 -4
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set.wgsl +109 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set_rows.wgsl +2 -3
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set_rows_quant.wgsl +224 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{soft_max.tmpl.wgsl → soft_max.wgsl} +106 -206
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/solve_tri.wgsl +121 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/ssm_conv.wgsl +65 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl +193 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/unary.wgsl +68 -48
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/upscale.wgsl +240 -0
- data/ext/sources/ggml/src/ggml-zdnn/ggml-zdnn.cpp +18 -14
- data/ext/sources/ggml/src/ggml-zendnn/CMakeLists.txt +1 -1
- data/ext/sources/ggml/src/ggml-zendnn/ggml-zendnn.cpp +244 -10
- data/ext/sources/ggml/src/ggml.c +146 -42
- data/ext/sources/ggml/src/gguf.cpp +173 -28
- data/ext/sources/include/parakeet.h +342 -0
- data/ext/sources/include/whisper.h +31 -0
- data/ext/sources/media/matmul.png +0 -0
- data/ext/sources/src/CMakeLists.txt +23 -0
- data/ext/sources/src/parakeet-arch.h +188 -0
- data/ext/sources/src/parakeet.cpp +3838 -0
- data/ext/sources/src/whisper.cpp +220 -26
- data/extsources.rb +26 -10
- data/lib/whisper/log_settable.rb +33 -0
- data/lib/whisper/model/uri.rb +13 -8
- data/lib/whisper/output.rb +74 -0
- data/sig/whisper.rbs +417 -62
- data/test/helper.rb +2 -0
- data/test/jfk_reader/jfk_reader.c +50 -7
- data/test/test_callback.rb +1 -0
- data/test/test_package.rb +6 -5
- data/test/test_parakeet.rb +28 -0
- data/test/test_parakeet_callback.rb +107 -0
- data/test/test_parakeet_context.rb +116 -0
- data/test/test_parakeet_context_params.rb +24 -0
- data/test/test_parakeet_model.rb +21 -0
- data/test/test_parakeet_params.rb +78 -0
- data/test/test_parakeet_segment.rb +42 -0
- data/test/test_parakeet_token.rb +73 -0
- data/test/test_params.rb +2 -0
- data/test/test_vad.rb +9 -0
- data/test/test_vad_context.rb +2 -2
- data/test/test_vad_segment.rb +1 -1
- data/test/test_whisper.rb +24 -6
- data/whispercpp.gemspec +2 -2
- metadata +263 -304
- data/ext/sources/bindings/javascript/CMakeLists.txt +0 -41
- data/ext/sources/bindings/javascript/emscripten.cpp +0 -93
- data/ext/sources/bindings/javascript/libwhisper.worker.js +0 -1
- data/ext/sources/bindings/javascript/package.json +0 -26
- data/ext/sources/bindings/javascript/whisper.js +0 -19
- data/ext/sources/examples/addon.node/CMakeLists.txt +0 -31
- data/ext/sources/examples/addon.node/__test__/whisper.spec.js +0 -133
- data/ext/sources/examples/addon.node/addon.cpp +0 -557
- data/ext/sources/examples/addon.node/index.js +0 -59
- data/ext/sources/examples/addon.node/package.json +0 -16
- data/ext/sources/examples/addon.node/vad-example.js +0 -132
- data/ext/sources/examples/bench.wasm/CMakeLists.txt +0 -49
- data/ext/sources/examples/bench.wasm/emscripten.cpp +0 -87
- data/ext/sources/examples/bench.wasm/index-tmpl.html +0 -285
- data/ext/sources/examples/coi-serviceworker.js +0 -146
- data/ext/sources/examples/command/CMakeLists.txt +0 -10
- data/ext/sources/examples/command/command.cpp +0 -802
- data/ext/sources/examples/command/commands.txt +0 -9
- data/ext/sources/examples/command.wasm/CMakeLists.txt +0 -50
- data/ext/sources/examples/command.wasm/emscripten.cpp +0 -327
- data/ext/sources/examples/command.wasm/index-tmpl.html +0 -415
- data/ext/sources/examples/generate-karaoke.sh +0 -57
- data/ext/sources/examples/helpers.js +0 -191
- data/ext/sources/examples/livestream.sh +0 -112
- data/ext/sources/examples/lsp/CMakeLists.txt +0 -10
- data/ext/sources/examples/lsp/lsp.cpp +0 -471
- data/ext/sources/examples/lsp/whisper.vim +0 -362
- data/ext/sources/examples/python/test_whisper_processor.py +0 -7
- data/ext/sources/examples/python/whisper_processor.py +0 -54
- data/ext/sources/examples/server/bench.js +0 -29
- data/ext/sources/examples/server.py +0 -120
- data/ext/sources/examples/stream/CMakeLists.txt +0 -10
- data/ext/sources/examples/stream/stream.cpp +0 -437
- data/ext/sources/examples/stream.wasm/CMakeLists.txt +0 -49
- data/ext/sources/examples/stream.wasm/emscripten.cpp +0 -216
- data/ext/sources/examples/stream.wasm/index-tmpl.html +0 -491
- data/ext/sources/examples/sycl/CMakeLists.txt +0 -9
- data/ext/sources/examples/sycl/build.sh +0 -22
- data/ext/sources/examples/sycl/ls-sycl-device.cpp +0 -11
- data/ext/sources/examples/sycl/run-whisper.sh +0 -17
- data/ext/sources/examples/talk-llama/CMakeLists.txt +0 -48
- data/ext/sources/examples/talk-llama/eleven-labs.py +0 -80
- data/ext/sources/examples/talk-llama/llama-adapter.cpp +0 -488
- data/ext/sources/examples/talk-llama/llama-adapter.h +0 -89
- data/ext/sources/examples/talk-llama/llama-arch.cpp +0 -2877
- data/ext/sources/examples/talk-llama/llama-arch.h +0 -628
- data/ext/sources/examples/talk-llama/llama-batch.cpp +0 -919
- data/ext/sources/examples/talk-llama/llama-batch.h +0 -173
- data/ext/sources/examples/talk-llama/llama-chat.cpp +0 -896
- data/ext/sources/examples/talk-llama/llama-chat.h +0 -71
- data/ext/sources/examples/talk-llama/llama-context.cpp +0 -3633
- data/ext/sources/examples/talk-llama/llama-context.h +0 -359
- data/ext/sources/examples/talk-llama/llama-cparams.cpp +0 -5
- data/ext/sources/examples/talk-llama/llama-cparams.h +0 -47
- data/ext/sources/examples/talk-llama/llama-ext.h +0 -12
- data/ext/sources/examples/talk-llama/llama-grammar.cpp +0 -1464
- data/ext/sources/examples/talk-llama/llama-grammar.h +0 -194
- data/ext/sources/examples/talk-llama/llama-graph.cpp +0 -2735
- data/ext/sources/examples/talk-llama/llama-graph.h +0 -1031
- data/ext/sources/examples/talk-llama/llama-hparams.cpp +0 -258
- data/ext/sources/examples/talk-llama/llama-hparams.h +0 -353
- data/ext/sources/examples/talk-llama/llama-impl.cpp +0 -171
- data/ext/sources/examples/talk-llama/llama-impl.h +0 -75
- data/ext/sources/examples/talk-llama/llama-io.cpp +0 -15
- data/ext/sources/examples/talk-llama/llama-io.h +0 -35
- data/ext/sources/examples/talk-llama/llama-kv-cache-iswa.cpp +0 -330
- data/ext/sources/examples/talk-llama/llama-kv-cache-iswa.h +0 -137
- data/ext/sources/examples/talk-llama/llama-kv-cache.cpp +0 -2285
- data/ext/sources/examples/talk-llama/llama-kv-cache.h +0 -389
- data/ext/sources/examples/talk-llama/llama-kv-cells.h +0 -533
- data/ext/sources/examples/talk-llama/llama-memory-hybrid-iswa.cpp +0 -275
- data/ext/sources/examples/talk-llama/llama-memory-hybrid-iswa.h +0 -140
- data/ext/sources/examples/talk-llama/llama-memory-hybrid.cpp +0 -268
- data/ext/sources/examples/talk-llama/llama-memory-hybrid.h +0 -139
- data/ext/sources/examples/talk-llama/llama-memory-recurrent.cpp +0 -1165
- data/ext/sources/examples/talk-llama/llama-memory-recurrent.h +0 -182
- data/ext/sources/examples/talk-llama/llama-memory.cpp +0 -59
- data/ext/sources/examples/talk-llama/llama-memory.h +0 -122
- data/ext/sources/examples/talk-llama/llama-mmap.cpp +0 -752
- data/ext/sources/examples/talk-llama/llama-mmap.h +0 -73
- data/ext/sources/examples/talk-llama/llama-model-loader.cpp +0 -1655
- data/ext/sources/examples/talk-llama/llama-model-loader.h +0 -206
- data/ext/sources/examples/talk-llama/llama-model-saver.cpp +0 -299
- data/ext/sources/examples/talk-llama/llama-model-saver.h +0 -40
- data/ext/sources/examples/talk-llama/llama-model.cpp +0 -9056
- data/ext/sources/examples/talk-llama/llama-model.h +0 -597
- data/ext/sources/examples/talk-llama/llama-quant.cpp +0 -1304
- data/ext/sources/examples/talk-llama/llama-quant.h +0 -1
- data/ext/sources/examples/talk-llama/llama-sampler.cpp +0 -3885
- data/ext/sources/examples/talk-llama/llama-sampler.h +0 -42
- data/ext/sources/examples/talk-llama/llama-vocab.cpp +0 -3970
- data/ext/sources/examples/talk-llama/llama-vocab.h +0 -187
- data/ext/sources/examples/talk-llama/llama.cpp +0 -1194
- data/ext/sources/examples/talk-llama/llama.h +0 -1573
- data/ext/sources/examples/talk-llama/models/afmoe.cpp +0 -190
- data/ext/sources/examples/talk-llama/models/apertus.cpp +0 -125
- data/ext/sources/examples/talk-llama/models/arcee.cpp +0 -135
- data/ext/sources/examples/talk-llama/models/arctic.cpp +0 -137
- data/ext/sources/examples/talk-llama/models/arwkv7.cpp +0 -86
- data/ext/sources/examples/talk-llama/models/baichuan.cpp +0 -123
- data/ext/sources/examples/talk-llama/models/bailingmoe.cpp +0 -143
- data/ext/sources/examples/talk-llama/models/bailingmoe2.cpp +0 -133
- data/ext/sources/examples/talk-llama/models/bert.cpp +0 -184
- data/ext/sources/examples/talk-llama/models/bitnet.cpp +0 -145
- data/ext/sources/examples/talk-llama/models/bloom.cpp +0 -101
- data/ext/sources/examples/talk-llama/models/chameleon.cpp +0 -178
- data/ext/sources/examples/talk-llama/models/chatglm.cpp +0 -132
- data/ext/sources/examples/talk-llama/models/codeshell.cpp +0 -111
- data/ext/sources/examples/talk-llama/models/cogvlm.cpp +0 -102
- data/ext/sources/examples/talk-llama/models/cohere2-iswa.cpp +0 -134
- data/ext/sources/examples/talk-llama/models/command-r.cpp +0 -122
- data/ext/sources/examples/talk-llama/models/dbrx.cpp +0 -122
- data/ext/sources/examples/talk-llama/models/deci.cpp +0 -135
- data/ext/sources/examples/talk-llama/models/deepseek.cpp +0 -142
- data/ext/sources/examples/talk-llama/models/deepseek2.cpp +0 -262
- data/ext/sources/examples/talk-llama/models/delta-net-base.cpp +0 -445
- data/ext/sources/examples/talk-llama/models/dots1.cpp +0 -132
- data/ext/sources/examples/talk-llama/models/dream.cpp +0 -105
- data/ext/sources/examples/talk-llama/models/ernie4-5-moe.cpp +0 -148
- data/ext/sources/examples/talk-llama/models/ernie4-5.cpp +0 -110
- data/ext/sources/examples/talk-llama/models/eurobert.cpp +0 -97
- data/ext/sources/examples/talk-llama/models/exaone-moe.cpp +0 -145
- data/ext/sources/examples/talk-llama/models/exaone.cpp +0 -114
- data/ext/sources/examples/talk-llama/models/exaone4.cpp +0 -123
- data/ext/sources/examples/talk-llama/models/falcon-h1.cpp +0 -111
- data/ext/sources/examples/talk-llama/models/falcon.cpp +0 -120
- data/ext/sources/examples/talk-llama/models/gemma-embedding.cpp +0 -116
- data/ext/sources/examples/talk-llama/models/gemma.cpp +0 -112
- data/ext/sources/examples/talk-llama/models/gemma2-iswa.cpp +0 -128
- data/ext/sources/examples/talk-llama/models/gemma3.cpp +0 -155
- data/ext/sources/examples/talk-llama/models/gemma3n-iswa.cpp +0 -384
- data/ext/sources/examples/talk-llama/models/glm4-moe.cpp +0 -170
- data/ext/sources/examples/talk-llama/models/glm4.cpp +0 -157
- data/ext/sources/examples/talk-llama/models/gpt2.cpp +0 -105
- data/ext/sources/examples/talk-llama/models/gptneox.cpp +0 -144
- data/ext/sources/examples/talk-llama/models/granite-hybrid.cpp +0 -195
- data/ext/sources/examples/talk-llama/models/granite.cpp +0 -210
- data/ext/sources/examples/talk-llama/models/grok.cpp +0 -159
- data/ext/sources/examples/talk-llama/models/grovemoe.cpp +0 -139
- data/ext/sources/examples/talk-llama/models/hunyuan-dense.cpp +0 -132
- data/ext/sources/examples/talk-llama/models/hunyuan-moe.cpp +0 -153
- data/ext/sources/examples/talk-llama/models/internlm2.cpp +0 -120
- data/ext/sources/examples/talk-llama/models/jais.cpp +0 -86
- data/ext/sources/examples/talk-llama/models/jais2.cpp +0 -123
- data/ext/sources/examples/talk-llama/models/jamba.cpp +0 -106
- data/ext/sources/examples/talk-llama/models/kimi-linear.cpp +0 -381
- data/ext/sources/examples/talk-llama/models/lfm2.cpp +0 -196
- data/ext/sources/examples/talk-llama/models/llada-moe.cpp +0 -122
- data/ext/sources/examples/talk-llama/models/llada.cpp +0 -99
- data/ext/sources/examples/talk-llama/models/llama-iswa.cpp +0 -178
- data/ext/sources/examples/talk-llama/models/llama.cpp +0 -175
- data/ext/sources/examples/talk-llama/models/maincoder.cpp +0 -117
- data/ext/sources/examples/talk-llama/models/mamba-base.cpp +0 -289
- data/ext/sources/examples/talk-llama/models/mamba.cpp +0 -54
- data/ext/sources/examples/talk-llama/models/mimo2-iswa.cpp +0 -129
- data/ext/sources/examples/talk-llama/models/minicpm3.cpp +0 -200
- data/ext/sources/examples/talk-llama/models/minimax-m2.cpp +0 -123
- data/ext/sources/examples/talk-llama/models/mistral3.cpp +0 -160
- data/ext/sources/examples/talk-llama/models/models.h +0 -704
- data/ext/sources/examples/talk-llama/models/modern-bert.cpp +0 -109
- data/ext/sources/examples/talk-llama/models/mpt.cpp +0 -126
- data/ext/sources/examples/talk-llama/models/nemotron-h.cpp +0 -162
- data/ext/sources/examples/talk-llama/models/nemotron.cpp +0 -122
- data/ext/sources/examples/talk-llama/models/neo-bert.cpp +0 -104
- data/ext/sources/examples/talk-llama/models/olmo.cpp +0 -121
- data/ext/sources/examples/talk-llama/models/olmo2.cpp +0 -150
- data/ext/sources/examples/talk-llama/models/olmoe.cpp +0 -124
- data/ext/sources/examples/talk-llama/models/openai-moe-iswa.cpp +0 -127
- data/ext/sources/examples/talk-llama/models/openelm.cpp +0 -124
- data/ext/sources/examples/talk-llama/models/orion.cpp +0 -123
- data/ext/sources/examples/talk-llama/models/paddleocr.cpp +0 -122
- data/ext/sources/examples/talk-llama/models/pangu-embedded.cpp +0 -121
- data/ext/sources/examples/talk-llama/models/phi2.cpp +0 -121
- data/ext/sources/examples/talk-llama/models/phi3.cpp +0 -152
- data/ext/sources/examples/talk-llama/models/plamo.cpp +0 -110
- data/ext/sources/examples/talk-llama/models/plamo2.cpp +0 -320
- data/ext/sources/examples/talk-llama/models/plamo3.cpp +0 -128
- data/ext/sources/examples/talk-llama/models/plm.cpp +0 -169
- data/ext/sources/examples/talk-llama/models/qwen.cpp +0 -108
- data/ext/sources/examples/talk-llama/models/qwen2.cpp +0 -126
- data/ext/sources/examples/talk-llama/models/qwen2moe.cpp +0 -151
- data/ext/sources/examples/talk-llama/models/qwen2vl.cpp +0 -117
- data/ext/sources/examples/talk-llama/models/qwen3.cpp +0 -120
- data/ext/sources/examples/talk-llama/models/qwen35.cpp +0 -381
- data/ext/sources/examples/talk-llama/models/qwen35moe.cpp +0 -422
- data/ext/sources/examples/talk-llama/models/qwen3moe.cpp +0 -131
- data/ext/sources/examples/talk-llama/models/qwen3next.cpp +0 -525
- data/ext/sources/examples/talk-llama/models/qwen3vl-moe.cpp +0 -140
- data/ext/sources/examples/talk-llama/models/qwen3vl.cpp +0 -132
- data/ext/sources/examples/talk-llama/models/refact.cpp +0 -94
- data/ext/sources/examples/talk-llama/models/rnd1.cpp +0 -126
- data/ext/sources/examples/talk-llama/models/rwkv6-base.cpp +0 -164
- data/ext/sources/examples/talk-llama/models/rwkv6.cpp +0 -94
- data/ext/sources/examples/talk-llama/models/rwkv6qwen2.cpp +0 -86
- data/ext/sources/examples/talk-llama/models/rwkv7-base.cpp +0 -137
- data/ext/sources/examples/talk-llama/models/rwkv7.cpp +0 -90
- data/ext/sources/examples/talk-llama/models/seed-oss.cpp +0 -124
- data/ext/sources/examples/talk-llama/models/smallthinker.cpp +0 -126
- data/ext/sources/examples/talk-llama/models/smollm3.cpp +0 -128
- data/ext/sources/examples/talk-llama/models/stablelm.cpp +0 -146
- data/ext/sources/examples/talk-llama/models/starcoder.cpp +0 -100
- data/ext/sources/examples/talk-llama/models/starcoder2.cpp +0 -121
- data/ext/sources/examples/talk-llama/models/step35-iswa.cpp +0 -165
- data/ext/sources/examples/talk-llama/models/t5-dec.cpp +0 -166
- data/ext/sources/examples/talk-llama/models/t5-enc.cpp +0 -96
- data/ext/sources/examples/talk-llama/models/wavtokenizer-dec.cpp +0 -149
- data/ext/sources/examples/talk-llama/models/xverse.cpp +0 -108
- data/ext/sources/examples/talk-llama/prompts/talk-alpaca.txt +0 -23
- data/ext/sources/examples/talk-llama/speak +0 -40
- data/ext/sources/examples/talk-llama/speak.bat +0 -1
- data/ext/sources/examples/talk-llama/speak.ps1 +0 -14
- data/ext/sources/examples/talk-llama/talk-llama.cpp +0 -813
- data/ext/sources/examples/talk-llama/unicode-data.cpp +0 -7034
- data/ext/sources/examples/talk-llama/unicode-data.h +0 -20
- data/ext/sources/examples/talk-llama/unicode.cpp +0 -1103
- data/ext/sources/examples/talk-llama/unicode.h +0 -111
- data/ext/sources/examples/wchess/CMakeLists.txt +0 -10
- data/ext/sources/examples/wchess/libwchess/CMakeLists.txt +0 -19
- data/ext/sources/examples/wchess/libwchess/Chessboard.cpp +0 -803
- data/ext/sources/examples/wchess/libwchess/Chessboard.h +0 -33
- data/ext/sources/examples/wchess/libwchess/WChess.cpp +0 -193
- data/ext/sources/examples/wchess/libwchess/WChess.h +0 -63
- data/ext/sources/examples/wchess/libwchess/test-chessboard.cpp +0 -117
- data/ext/sources/examples/wchess/wchess.cmd/CMakeLists.txt +0 -8
- data/ext/sources/examples/wchess/wchess.cmd/wchess.cmd.cpp +0 -253
- data/ext/sources/examples/whisper.wasm/CMakeLists.txt +0 -50
- data/ext/sources/examples/whisper.wasm/emscripten.cpp +0 -118
- data/ext/sources/examples/whisper.wasm/index-tmpl.html +0 -659
- data/ext/sources/ggml/src/ggml-cuda/template-instances/generate_cu_files.py +0 -99
- data/ext/sources/ggml/src/ggml-hexagon/htp/htp-msg.h +0 -155
- data/ext/sources/ggml/src/ggml-hexagon/op-desc.h +0 -153
- data/ext/sources/ggml/src/ggml-opencl/kernels/embed_kernel.py +0 -26
- data/ext/sources/ggml/src/ggml-openvino/openvino/pass/eliminate_zp.cpp +0 -123
- data/ext/sources/ggml/src/ggml-openvino/openvino/pass/eliminate_zp.h +0 -17
- data/ext/sources/ggml/src/ggml-virtgpu/regenerate_remoting.py +0 -333
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/abs.comp +0 -21
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ceil.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/clamp.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/cos.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/elu.comp +0 -27
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/exp.comp +0 -21
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/floor.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu.comp +0 -25
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_erf.comp +0 -39
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_quick.comp +0 -23
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardsigmoid.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardswish.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/leaky_relu.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/neg.comp +0 -20
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/relu.comp +0 -21
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/round.comp +0 -29
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rte.glsl +0 -5
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sgn.comp +0 -21
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sigmoid.comp +0 -20
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/silu.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sin.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/softplus.comp +0 -23
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sqrt.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/square.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/step.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tanh.comp +0 -20
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/trunc.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/xielu.comp +0 -35
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/embed_wgsl.py +0 -182
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/glu.tmpl.wgsl +0 -323
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat.wgsl +0 -718
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm.wgsl +0 -123
- data/ext/sources/tests/CMakeLists.txt +0 -112
- data/ext/sources/tests/earnings21/eval.mk +0 -58
- data/ext/sources/tests/earnings21/eval.py +0 -68
- data/ext/sources/tests/earnings21/normalizers/__init__.py +0 -2
- data/ext/sources/tests/earnings21/normalizers/basic.py +0 -80
- data/ext/sources/tests/earnings21/normalizers/english.json +0 -1741
- data/ext/sources/tests/earnings21/normalizers/english.py +0 -550
- data/ext/sources/tests/earnings21/requirements.txt +0 -6
- data/ext/sources/tests/en-0-ref.txt +0 -1
- data/ext/sources/tests/en-1-ref.txt +0 -1
- data/ext/sources/tests/en-2-ref.txt +0 -1
- data/ext/sources/tests/es-0-ref.txt +0 -1
- data/ext/sources/tests/librispeech/eval.mk +0 -39
- data/ext/sources/tests/librispeech/eval.py +0 -47
- data/ext/sources/tests/librispeech/normalizers/__init__.py +0 -2
- data/ext/sources/tests/librispeech/normalizers/basic.py +0 -80
- data/ext/sources/tests/librispeech/normalizers/english.json +0 -1741
- data/ext/sources/tests/librispeech/normalizers/english.py +0 -550
- data/ext/sources/tests/librispeech/requirements.txt +0 -6
- data/ext/sources/tests/run-tests.sh +0 -130
- data/ext/sources/tests/test-c.c +0 -3
- data/ext/sources/tests/test-vad-full.cpp +0 -56
- data/ext/sources/tests/test-vad.cpp +0 -83
- data/ext/sources/tests/test-whisper.js +0 -58
- data/lib/whisper/context.rb +0 -15
- data/lib/whisper/segment.rb +0 -58
|
@@ -56,6 +56,65 @@ static void mul_mat_vec_q_reorder(const void * __restrict__ vx, const void * __r
|
|
|
56
56
|
}
|
|
57
57
|
}
|
|
58
58
|
|
|
59
|
+
template <typename reorder_vec_dot_q_sycl, int ncols_dst>
|
|
60
|
+
static void mul_mat_vec_q_reorder_ncols(const void * __restrict__ vx, const void * __restrict__ vy,
|
|
61
|
+
float * __restrict__ dst, const int ncols, const int nrows,
|
|
62
|
+
const int stride_col_y_bytes, const int stride_col_dst,
|
|
63
|
+
const sycl::nd_item<3> & nd_item) {
|
|
64
|
+
using block_type = ggml_sycl_reordered::block_q_t<reorder_vec_dot_q_sycl::gtype>;
|
|
65
|
+
using block_traits = typename block_type::traits;
|
|
66
|
+
|
|
67
|
+
const auto sg = nd_item.get_sub_group();
|
|
68
|
+
const int sg_range = sg.get_group_linear_range();
|
|
69
|
+
const int workgroup_id = nd_item.get_group_linear_id();
|
|
70
|
+
const int sg_id = sg.get_group_linear_id();
|
|
71
|
+
const int row = workgroup_id * sg_range + sg_id;
|
|
72
|
+
|
|
73
|
+
if (row >= nrows) {
|
|
74
|
+
return;
|
|
75
|
+
}
|
|
76
|
+
|
|
77
|
+
const int blocks_per_row = ncols / block_traits::qk;
|
|
78
|
+
constexpr int blocks_per_subgroup = ceil_div(block_traits::vdr_mmvq * WARP_SIZE, block_traits::qi);
|
|
79
|
+
constexpr int block_elements_per_subgroup = block_traits::qi / block_traits::vdr_mmvq;
|
|
80
|
+
const int nblocks = nrows * (ncols / block_traits::qk);
|
|
81
|
+
|
|
82
|
+
static_assert(blocks_per_subgroup > 0);
|
|
83
|
+
static_assert(block_elements_per_subgroup > 0);
|
|
84
|
+
|
|
85
|
+
float partial_sum[ncols_dst] = {0.0f};
|
|
86
|
+
for (int i = sg.get_local_linear_id() / block_elements_per_subgroup; i < blocks_per_row; i += blocks_per_subgroup) {
|
|
87
|
+
const int ibx = row * blocks_per_row + i;
|
|
88
|
+
|
|
89
|
+
const auto bx_offset = block_type::get_block_offset(ibx, nblocks);
|
|
90
|
+
const auto d_offset = block_type::get_d_offset(nrows, ncols, ibx);
|
|
91
|
+
const int iby = i * block_type::block_to_q8_1_ratio();
|
|
92
|
+
|
|
93
|
+
#pragma unroll
|
|
94
|
+
for (int elem = 0; elem < block_elements_per_subgroup; elem += WARP_SIZE) {
|
|
95
|
+
const int iqs = elem + block_traits::vdr_mmvq * (sg.get_local_linear_id() % block_elements_per_subgroup);
|
|
96
|
+
|
|
97
|
+
#pragma unroll
|
|
98
|
+
for (int j = 0; j < ncols_dst; ++j) {
|
|
99
|
+
const char * vy_j = (const char *)vy + j * stride_col_y_bytes;
|
|
100
|
+
const int8_t * q8_1_quant_ptr = (const int8_t *)vy_j + iby * QK8_1;
|
|
101
|
+
const sycl::half2* q8_1_ds_ptr = (const sycl::half2 *)(vy_j + ncols + iby * sizeof(sycl::half2));
|
|
102
|
+
|
|
103
|
+
partial_sum[j] += reorder_vec_dot_q_sycl()(vx, bx_offset, d_offset, q8_1_quant_ptr, q8_1_ds_ptr, iqs);
|
|
104
|
+
}
|
|
105
|
+
}
|
|
106
|
+
}
|
|
107
|
+
|
|
108
|
+
#pragma unroll
|
|
109
|
+
for (int j = 0; j < ncols_dst; ++j) {
|
|
110
|
+
float sum = sycl::reduce_over_group(nd_item.get_sub_group(), partial_sum[j], std::plus<>());
|
|
111
|
+
|
|
112
|
+
if (sg.leader()) {
|
|
113
|
+
dst[j * stride_col_dst + row] = sum;
|
|
114
|
+
}
|
|
115
|
+
}
|
|
116
|
+
}
|
|
117
|
+
|
|
59
118
|
template <int qk, int qi, typename block_q_t, int vdr, vec_dot_q_sycl_t vec_dot_q_sycl>
|
|
60
119
|
static void mul_mat_vec_q(const void * __restrict__ vx, const void * __restrict__ vy, float * __restrict__ dst,
|
|
61
120
|
const int ncols, const int nrows, const sycl::nd_item<3> & item_ct1) {
|
|
@@ -100,6 +159,70 @@ static void mul_mat_vec_q(const void * __restrict__ vx, const void * __restrict_
|
|
|
100
159
|
}
|
|
101
160
|
}
|
|
102
161
|
|
|
162
|
+
template <int qk, int qi, typename block_q_t, int vdr,
|
|
163
|
+
vec_dot_q_sycl_t vec_dot_q_sycl, int ncols_dst>
|
|
164
|
+
static void mul_mat_vec_q_ncols(
|
|
165
|
+
const void * __restrict__ vx,
|
|
166
|
+
const void * __restrict__ vy,
|
|
167
|
+
float * __restrict__ dst,
|
|
168
|
+
const int ncols,
|
|
169
|
+
const int nrows,
|
|
170
|
+
const int stride_col_y,
|
|
171
|
+
const int stride_col_dst,
|
|
172
|
+
const sycl::nd_item<3> & item_ct1) {
|
|
173
|
+
|
|
174
|
+
const int row = item_ct1.get_group(2) * item_ct1.get_local_range(1)
|
|
175
|
+
+ item_ct1.get_local_id(1);
|
|
176
|
+
|
|
177
|
+
if (row >= nrows) {
|
|
178
|
+
return;
|
|
179
|
+
}
|
|
180
|
+
|
|
181
|
+
const int blocks_per_row = ncols / qk;
|
|
182
|
+
constexpr int blocks_per_warp = (vdr * WARP_SIZE + qi - 1) / qi;
|
|
183
|
+
|
|
184
|
+
// partial sums: one per output column
|
|
185
|
+
float tmp[ncols_dst] = {0.0f};
|
|
186
|
+
|
|
187
|
+
const block_q_t * x = (const block_q_t *) vx;
|
|
188
|
+
const block_q8_1 * y = (const block_q8_1 *) vy;
|
|
189
|
+
|
|
190
|
+
for (int i = item_ct1.get_local_id(2) / (qi / vdr);
|
|
191
|
+
i < blocks_per_row;
|
|
192
|
+
i += blocks_per_warp) {
|
|
193
|
+
|
|
194
|
+
const int ibx = row * blocks_per_row + i;
|
|
195
|
+
const int iby = i * (qk / QK8_1);
|
|
196
|
+
|
|
197
|
+
// read weight block once, dot against all columns
|
|
198
|
+
for (size_t elem = 0; elem < qi / vdr; elem += WARP_SIZE) {
|
|
199
|
+
const int iqs = elem + vdr * (item_ct1.get_local_id(2) % (qi / vdr));
|
|
200
|
+
|
|
201
|
+
#pragma unroll
|
|
202
|
+
for (int j = 0; j < ncols_dst; ++j) {
|
|
203
|
+
tmp[j] += vec_dot_q_sycl(&x[ibx], &y[j * stride_col_y + iby], iqs);
|
|
204
|
+
}
|
|
205
|
+
}
|
|
206
|
+
}
|
|
207
|
+
|
|
208
|
+
// reduce within subgroup
|
|
209
|
+
#pragma unroll
|
|
210
|
+
for (int j = 0; j < ncols_dst; ++j) {
|
|
211
|
+
#pragma unroll
|
|
212
|
+
for (int mask = WARP_SIZE / 2; mask > 0; mask >>= 1) {
|
|
213
|
+
tmp[j] += dpct::permute_sub_group_by_xor(
|
|
214
|
+
item_ct1.get_sub_group(), tmp[j], mask);
|
|
215
|
+
}
|
|
216
|
+
}
|
|
217
|
+
|
|
218
|
+
if (item_ct1.get_local_id(2) == 0) {
|
|
219
|
+
#pragma unroll
|
|
220
|
+
for (int j = 0; j < ncols_dst; ++j) {
|
|
221
|
+
dst[j * stride_col_dst + row] = tmp[j];
|
|
222
|
+
}
|
|
223
|
+
}
|
|
224
|
+
}
|
|
225
|
+
|
|
103
226
|
template <int qk, int qi, typename block_q_t, int vdr>
|
|
104
227
|
static void mul_mat_vec_q_iq2_xxs_q8_1(const void *__restrict__ vx,
|
|
105
228
|
const void *__restrict__ vy,
|
|
@@ -537,15 +660,14 @@ static void mul_mat_vec_q_iq4_xs_q8_1(const void *__restrict__ vx,
|
|
|
537
660
|
static void reorder_mul_mat_vec_q4_0_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols,
|
|
538
661
|
const int nrows, dpct::queue_ptr stream) {
|
|
539
662
|
GGML_ASSERT(ncols % QK4_0 == 0);
|
|
540
|
-
|
|
541
|
-
constexpr size_t num_subgroups =
|
|
542
|
-
|
|
543
|
-
|
|
544
|
-
const sycl::range<3>
|
|
545
|
-
const sycl::range<3> workgroup_size(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
663
|
+
// Round up to a whole number of subgroup-sized workgroups; out-of-range rows are skipped inside the kernel.
|
|
664
|
+
constexpr size_t num_subgroups = WARP_SIZE;
|
|
665
|
+
const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
|
|
666
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
667
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
546
668
|
|
|
547
669
|
stream->submit([&](sycl::handler & cgh) {
|
|
548
|
-
cgh.parallel_for(sycl::nd_range<3>(
|
|
670
|
+
cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
549
671
|
[=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
550
672
|
mul_mat_vec_q_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q4_0>>(vx, vy, dst, ncols, nrows,
|
|
551
673
|
nd_item);
|
|
@@ -553,6 +675,45 @@ static void reorder_mul_mat_vec_q4_0_q8_1_sycl(const void * vx, const void * vy,
|
|
|
553
675
|
});
|
|
554
676
|
}
|
|
555
677
|
|
|
678
|
+
template <int ncols_dst>
|
|
679
|
+
static void reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols(
|
|
680
|
+
const void * vx, const void * vy, float * dst,
|
|
681
|
+
const int ncols, const int nrows,
|
|
682
|
+
const int stride_col_y_bytes, const int stride_col_dst,
|
|
683
|
+
dpct::queue_ptr stream) {
|
|
684
|
+
GGML_ASSERT(ncols % QK4_0 == 0);
|
|
685
|
+
constexpr size_t num_subgroups = WARP_SIZE;
|
|
686
|
+
const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
|
|
687
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
688
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
689
|
+
|
|
690
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
691
|
+
cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
692
|
+
[=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
693
|
+
mul_mat_vec_q_reorder_ncols<reorder_vec_dot_q_sycl<GGML_TYPE_Q4_0>, ncols_dst>(
|
|
694
|
+
vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, nd_item);
|
|
695
|
+
});
|
|
696
|
+
});
|
|
697
|
+
}
|
|
698
|
+
|
|
699
|
+
static void reorder_mul_mat_vec_q4_0_q8_1_sycl_switch_ncols(
|
|
700
|
+
const void * vx, const void * vy, float * dst,
|
|
701
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
702
|
+
const int stride_col_y_bytes, const int stride_col_dst,
|
|
703
|
+
dpct::queue_ptr stream) {
|
|
704
|
+
switch (ncols_dst) {
|
|
705
|
+
case 1: reorder_mul_mat_vec_q4_0_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
706
|
+
case 2: reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
707
|
+
case 3: reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
708
|
+
case 4: reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
709
|
+
case 5: reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
710
|
+
case 6: reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
711
|
+
case 7: reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
712
|
+
case 8: reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
713
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q4_0 reorder multi-col MMVQ", ncols_dst);
|
|
714
|
+
}
|
|
715
|
+
}
|
|
716
|
+
|
|
556
717
|
static void mul_mat_vec_q4_0_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols, const int nrows,
|
|
557
718
|
dpct::queue_ptr stream) {
|
|
558
719
|
GGML_ASSERT(ncols % QK4_0 == 0);
|
|
@@ -571,6 +732,45 @@ static void mul_mat_vec_q4_0_q8_1_sycl(const void * vx, const void * vy, float *
|
|
|
571
732
|
}
|
|
572
733
|
}
|
|
573
734
|
|
|
735
|
+
template <int ncols_dst>
|
|
736
|
+
static void mul_mat_vec_q4_0_q8_1_sycl_ncols(
|
|
737
|
+
const void * vx, const void * vy, float * dst,
|
|
738
|
+
const int ncols, const int nrows,
|
|
739
|
+
const int stride_col_y, const int stride_col_dst,
|
|
740
|
+
dpct::queue_ptr stream) {
|
|
741
|
+
GGML_ASSERT(ncols % QK4_0 == 0);
|
|
742
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
743
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
744
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
745
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
746
|
+
cgh.parallel_for(
|
|
747
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
748
|
+
[=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
749
|
+
mul_mat_vec_q_ncols<QK4_0, QI4_0, block_q4_0,
|
|
750
|
+
VDR_Q4_0_Q8_1_MMVQ, vec_dot_q4_0_q8_1, ncols_dst>(
|
|
751
|
+
vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
|
|
752
|
+
});
|
|
753
|
+
});
|
|
754
|
+
}
|
|
755
|
+
|
|
756
|
+
static void mul_mat_vec_q4_0_q8_1_sycl_switch_ncols(
|
|
757
|
+
const void * vx, const void * vy, float * dst,
|
|
758
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
759
|
+
const int stride_col_y, const int stride_col_dst,
|
|
760
|
+
dpct::queue_ptr stream) {
|
|
761
|
+
switch (ncols_dst) {
|
|
762
|
+
case 1: mul_mat_vec_q4_0_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
763
|
+
case 2: mul_mat_vec_q4_0_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
764
|
+
case 3: mul_mat_vec_q4_0_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
765
|
+
case 4: mul_mat_vec_q4_0_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
766
|
+
case 5: mul_mat_vec_q4_0_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
767
|
+
case 6: mul_mat_vec_q4_0_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
768
|
+
case 7: mul_mat_vec_q4_0_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
769
|
+
case 8: mul_mat_vec_q4_0_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
770
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q4_0 multi-col MMVQ", ncols_dst);
|
|
771
|
+
}
|
|
772
|
+
}
|
|
773
|
+
|
|
574
774
|
static void mul_mat_vec_q4_1_q8_1_sycl(const void *vx, const void *vy,
|
|
575
775
|
float *dst, const int ncols,
|
|
576
776
|
const int nrows,
|
|
@@ -595,6 +795,45 @@ static void mul_mat_vec_q4_1_q8_1_sycl(const void *vx, const void *vy,
|
|
|
595
795
|
}
|
|
596
796
|
}
|
|
597
797
|
|
|
798
|
+
template <int ncols_dst>
|
|
799
|
+
static void mul_mat_vec_q4_1_q8_1_sycl_ncols(
|
|
800
|
+
const void * vx, const void * vy, float * dst,
|
|
801
|
+
const int ncols, const int nrows,
|
|
802
|
+
const int stride_col_y, const int stride_col_dst,
|
|
803
|
+
dpct::queue_ptr stream) {
|
|
804
|
+
GGML_ASSERT(ncols % QK4_1 == 0);
|
|
805
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
806
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
807
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
808
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
809
|
+
cgh.parallel_for(
|
|
810
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
811
|
+
[=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
812
|
+
mul_mat_vec_q_ncols<QK4_0, QI4_1, block_q4_1,
|
|
813
|
+
VDR_Q4_1_Q8_1_MMVQ, vec_dot_q4_1_q8_1, ncols_dst>(
|
|
814
|
+
vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
|
|
815
|
+
});
|
|
816
|
+
});
|
|
817
|
+
}
|
|
818
|
+
|
|
819
|
+
static void mul_mat_vec_q4_1_q8_1_sycl_switch_ncols(
|
|
820
|
+
const void * vx, const void * vy, float * dst,
|
|
821
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
822
|
+
const int stride_col_y, const int stride_col_dst,
|
|
823
|
+
dpct::queue_ptr stream) {
|
|
824
|
+
switch (ncols_dst) {
|
|
825
|
+
case 1: mul_mat_vec_q4_1_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
826
|
+
case 2: mul_mat_vec_q4_1_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
827
|
+
case 3: mul_mat_vec_q4_1_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
828
|
+
case 4: mul_mat_vec_q4_1_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
829
|
+
case 5: mul_mat_vec_q4_1_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
830
|
+
case 6: mul_mat_vec_q4_1_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
831
|
+
case 7: mul_mat_vec_q4_1_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
832
|
+
case 8: mul_mat_vec_q4_1_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
833
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q4_1 multi-col MMVQ", ncols_dst);
|
|
834
|
+
}
|
|
835
|
+
}
|
|
836
|
+
|
|
598
837
|
static void mul_mat_vec_mxfp4_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols, const int nrows,
|
|
599
838
|
dpct::queue_ptr stream) {
|
|
600
839
|
GGML_ASSERT(ncols % QK_MXFP4 == 0);
|
|
@@ -613,6 +852,101 @@ static void mul_mat_vec_mxfp4_q8_1_sycl(const void * vx, const void * vy, float
|
|
|
613
852
|
}
|
|
614
853
|
}
|
|
615
854
|
|
|
855
|
+
template <int ncols_dst>
|
|
856
|
+
static void mul_mat_vec_mxfp4_q8_1_sycl_ncols(
|
|
857
|
+
const void * vx, const void * vy, float * dst,
|
|
858
|
+
const int ncols, const int nrows,
|
|
859
|
+
const int stride_col_y, const int stride_col_dst,
|
|
860
|
+
dpct::queue_ptr stream) {
|
|
861
|
+
GGML_ASSERT(ncols % QK_MXFP4 == 0);
|
|
862
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
863
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
864
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
865
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
866
|
+
cgh.parallel_for(
|
|
867
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
868
|
+
[=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
869
|
+
mul_mat_vec_q_ncols<QK_MXFP4, QI_MXFP4, block_mxfp4,
|
|
870
|
+
VDR_MXFP4_Q8_1_MMVQ, vec_dot_mxfp4_q8_1, ncols_dst>(
|
|
871
|
+
vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
|
|
872
|
+
});
|
|
873
|
+
});
|
|
874
|
+
}
|
|
875
|
+
|
|
876
|
+
static void mul_mat_vec_mxfp4_q8_1_sycl_switch_ncols(
|
|
877
|
+
const void * vx, const void * vy, float * dst,
|
|
878
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
879
|
+
const int stride_col_y, const int stride_col_dst,
|
|
880
|
+
dpct::queue_ptr stream) {
|
|
881
|
+
switch (ncols_dst) {
|
|
882
|
+
case 1: mul_mat_vec_mxfp4_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
883
|
+
case 2: mul_mat_vec_mxfp4_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
884
|
+
case 3: mul_mat_vec_mxfp4_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
885
|
+
case 4: mul_mat_vec_mxfp4_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
886
|
+
case 5: mul_mat_vec_mxfp4_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
887
|
+
case 6: mul_mat_vec_mxfp4_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
888
|
+
case 7: mul_mat_vec_mxfp4_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
889
|
+
case 8: mul_mat_vec_mxfp4_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
890
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for MXFP4 multi-col MMVQ", ncols_dst);
|
|
891
|
+
}
|
|
892
|
+
}
|
|
893
|
+
|
|
894
|
+
static void mul_mat_vec_nvfp4_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols, const int nrows,
|
|
895
|
+
dpct::queue_ptr stream) {
|
|
896
|
+
GGML_ASSERT(ncols % QK_NVFP4 == 0);
|
|
897
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
898
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
899
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
900
|
+
|
|
901
|
+
{
|
|
902
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
903
|
+
cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
904
|
+
[=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
905
|
+
mul_mat_vec_q<QK_NVFP4, QI_NVFP4, block_nvfp4, VDR_NVFP4_Q8_1_MMVQ, vec_dot_nvfp4_q8_1>(
|
|
906
|
+
vx, vy, dst, ncols, nrows, item_ct1);
|
|
907
|
+
});
|
|
908
|
+
});
|
|
909
|
+
}
|
|
910
|
+
}
|
|
911
|
+
|
|
912
|
+
template <int ncols_dst>
|
|
913
|
+
static void mul_mat_vec_nvfp4_q8_1_sycl_ncols(
|
|
914
|
+
const void * vx, const void * vy, float * dst,
|
|
915
|
+
const int ncols, const int nrows,
|
|
916
|
+
const int stride_col_y, const int stride_col_dst,
|
|
917
|
+
dpct::queue_ptr stream) {
|
|
918
|
+
GGML_ASSERT(ncols % QK_NVFP4 == 0);
|
|
919
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
920
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
921
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
922
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
923
|
+
cgh.parallel_for(
|
|
924
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
925
|
+
[=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
926
|
+
mul_mat_vec_q_ncols<QK_NVFP4, QI_NVFP4, block_nvfp4,
|
|
927
|
+
VDR_NVFP4_Q8_1_MMVQ, vec_dot_nvfp4_q8_1, ncols_dst>(
|
|
928
|
+
vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
|
|
929
|
+
});
|
|
930
|
+
});
|
|
931
|
+
}
|
|
932
|
+
|
|
933
|
+
static void mul_mat_vec_nvfp4_q8_1_sycl_switch_ncols(
|
|
934
|
+
const void * vx, const void * vy, float * dst,
|
|
935
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
936
|
+
const int stride_col_y, const int stride_col_dst,
|
|
937
|
+
dpct::queue_ptr stream) {
|
|
938
|
+
switch (ncols_dst) {
|
|
939
|
+
case 1: mul_mat_vec_nvfp4_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
940
|
+
case 2: mul_mat_vec_nvfp4_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
941
|
+
case 3: mul_mat_vec_nvfp4_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
942
|
+
case 4: mul_mat_vec_nvfp4_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
943
|
+
case 5: mul_mat_vec_nvfp4_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
944
|
+
case 6: mul_mat_vec_nvfp4_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
945
|
+
case 7: mul_mat_vec_nvfp4_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
946
|
+
case 8: mul_mat_vec_nvfp4_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
947
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for NVFP4 multi-col MMVQ", ncols_dst);
|
|
948
|
+
}
|
|
949
|
+
}
|
|
616
950
|
|
|
617
951
|
static void mul_mat_vec_q5_0_q8_1_sycl(const void *vx, const void *vy,
|
|
618
952
|
float *dst, const int ncols,
|
|
@@ -638,6 +972,45 @@ static void mul_mat_vec_q5_0_q8_1_sycl(const void *vx, const void *vy,
|
|
|
638
972
|
}
|
|
639
973
|
}
|
|
640
974
|
|
|
975
|
+
template <int ncols_dst>
|
|
976
|
+
static void mul_mat_vec_q5_0_q8_1_sycl_ncols(
|
|
977
|
+
const void * vx, const void * vy, float * dst,
|
|
978
|
+
const int ncols, const int nrows,
|
|
979
|
+
const int stride_col_y, const int stride_col_dst,
|
|
980
|
+
dpct::queue_ptr stream) {
|
|
981
|
+
GGML_ASSERT(ncols % QK5_0 == 0);
|
|
982
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
983
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
984
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
985
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
986
|
+
cgh.parallel_for(
|
|
987
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
988
|
+
[=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
989
|
+
mul_mat_vec_q_ncols<QK5_0, QI5_0, block_q5_0,
|
|
990
|
+
VDR_Q5_0_Q8_1_MMVQ, vec_dot_q5_0_q8_1, ncols_dst>(
|
|
991
|
+
vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
|
|
992
|
+
});
|
|
993
|
+
});
|
|
994
|
+
}
|
|
995
|
+
|
|
996
|
+
static void mul_mat_vec_q5_0_q8_1_sycl_switch_ncols(
|
|
997
|
+
const void * vx, const void * vy, float * dst,
|
|
998
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
999
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1000
|
+
dpct::queue_ptr stream) {
|
|
1001
|
+
switch (ncols_dst) {
|
|
1002
|
+
case 1: mul_mat_vec_q5_0_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
1003
|
+
case 2: mul_mat_vec_q5_0_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1004
|
+
case 3: mul_mat_vec_q5_0_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1005
|
+
case 4: mul_mat_vec_q5_0_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1006
|
+
case 5: mul_mat_vec_q5_0_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1007
|
+
case 6: mul_mat_vec_q5_0_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1008
|
+
case 7: mul_mat_vec_q5_0_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1009
|
+
case 8: mul_mat_vec_q5_0_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1010
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q5_0 multi-col MMVQ", ncols_dst);
|
|
1011
|
+
}
|
|
1012
|
+
}
|
|
1013
|
+
|
|
641
1014
|
static void mul_mat_vec_q5_1_q8_1_sycl(const void *vx, const void *vy,
|
|
642
1015
|
float *dst, const int ncols,
|
|
643
1016
|
const int nrows,
|
|
@@ -662,6 +1035,102 @@ static void mul_mat_vec_q5_1_q8_1_sycl(const void *vx, const void *vy,
|
|
|
662
1035
|
}
|
|
663
1036
|
}
|
|
664
1037
|
|
|
1038
|
+
template <int ncols_dst>
|
|
1039
|
+
static void mul_mat_vec_q5_1_q8_1_sycl_ncols(
|
|
1040
|
+
const void * vx, const void * vy, float * dst,
|
|
1041
|
+
const int ncols, const int nrows,
|
|
1042
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1043
|
+
dpct::queue_ptr stream) {
|
|
1044
|
+
GGML_ASSERT(ncols % QK5_1 == 0);
|
|
1045
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
1046
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1047
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
1048
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1049
|
+
cgh.parallel_for(
|
|
1050
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1051
|
+
[=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1052
|
+
mul_mat_vec_q_ncols<QK5_1, QI5_1, block_q5_1,
|
|
1053
|
+
VDR_Q5_1_Q8_1_MMVQ, vec_dot_q5_1_q8_1, ncols_dst>(
|
|
1054
|
+
vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
|
|
1055
|
+
});
|
|
1056
|
+
});
|
|
1057
|
+
}
|
|
1058
|
+
|
|
1059
|
+
static void mul_mat_vec_q5_1_q8_1_sycl_switch_ncols(
|
|
1060
|
+
const void * vx, const void * vy, float * dst,
|
|
1061
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
1062
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1063
|
+
dpct::queue_ptr stream) {
|
|
1064
|
+
switch (ncols_dst) {
|
|
1065
|
+
case 1: mul_mat_vec_q5_1_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
1066
|
+
case 2: mul_mat_vec_q5_1_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1067
|
+
case 3: mul_mat_vec_q5_1_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1068
|
+
case 4: mul_mat_vec_q5_1_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1069
|
+
case 5: mul_mat_vec_q5_1_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1070
|
+
case 6: mul_mat_vec_q5_1_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1071
|
+
case 7: mul_mat_vec_q5_1_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1072
|
+
case 8: mul_mat_vec_q5_1_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1073
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q5_1 multi-col MMVQ", ncols_dst);
|
|
1074
|
+
}
|
|
1075
|
+
}
|
|
1076
|
+
|
|
1077
|
+
static void reorder_mul_mat_vec_q8_0_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols,
|
|
1078
|
+
const int nrows, dpct::queue_ptr stream) {
|
|
1079
|
+
GGML_ASSERT(ncols % QK8_0 == 0);
|
|
1080
|
+
// Round up to a whole number of subgroup-sized workgroups; out-of-range rows are skipped inside the kernel.
|
|
1081
|
+
constexpr size_t num_subgroups = WARP_SIZE;
|
|
1082
|
+
const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
|
|
1083
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1084
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
1085
|
+
|
|
1086
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1087
|
+
cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1088
|
+
[=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1089
|
+
mul_mat_vec_q_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q8_0>>(vx, vy, dst, ncols, nrows,
|
|
1090
|
+
nd_item);
|
|
1091
|
+
});
|
|
1092
|
+
});
|
|
1093
|
+
}
|
|
1094
|
+
|
|
1095
|
+
template <int ncols_dst>
|
|
1096
|
+
static void reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols(
|
|
1097
|
+
const void * vx, const void * vy, float * dst,
|
|
1098
|
+
const int ncols, const int nrows,
|
|
1099
|
+
const int stride_col_y_bytes, const int stride_col_dst,
|
|
1100
|
+
dpct::queue_ptr stream) {
|
|
1101
|
+
GGML_ASSERT(ncols % QK8_0 == 0);
|
|
1102
|
+
constexpr size_t num_subgroups = WARP_SIZE;
|
|
1103
|
+
const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
|
|
1104
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1105
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
1106
|
+
|
|
1107
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1108
|
+
cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1109
|
+
[=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1110
|
+
mul_mat_vec_q_reorder_ncols<reorder_vec_dot_q_sycl<GGML_TYPE_Q8_0>, ncols_dst>(
|
|
1111
|
+
vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, nd_item);
|
|
1112
|
+
});
|
|
1113
|
+
});
|
|
1114
|
+
}
|
|
1115
|
+
|
|
1116
|
+
static void reorder_mul_mat_vec_q8_0_q8_1_sycl_switch_ncols(
|
|
1117
|
+
const void * vx, const void * vy, float * dst,
|
|
1118
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
1119
|
+
const int stride_col_y_bytes, const int stride_col_dst,
|
|
1120
|
+
dpct::queue_ptr stream) {
|
|
1121
|
+
switch (ncols_dst) {
|
|
1122
|
+
case 1: reorder_mul_mat_vec_q8_0_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
1123
|
+
case 2: reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1124
|
+
case 3: reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1125
|
+
case 4: reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1126
|
+
case 5: reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1127
|
+
case 6: reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1128
|
+
case 7: reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1129
|
+
case 8: reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1130
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q8_0 reorder multi-col MMVQ", ncols_dst);
|
|
1131
|
+
}
|
|
1132
|
+
}
|
|
1133
|
+
|
|
665
1134
|
static void mul_mat_vec_q8_0_q8_1_sycl(const void *vx, const void *vy,
|
|
666
1135
|
float *dst, const int ncols,
|
|
667
1136
|
const int nrows,
|
|
@@ -686,6 +1155,105 @@ static void mul_mat_vec_q8_0_q8_1_sycl(const void *vx, const void *vy,
|
|
|
686
1155
|
}
|
|
687
1156
|
}
|
|
688
1157
|
|
|
1158
|
+
template <int ncols_dst>
|
|
1159
|
+
static void mul_mat_vec_q8_0_q8_1_sycl_ncols(
|
|
1160
|
+
const void * vx, const void * vy, float * dst,
|
|
1161
|
+
const int ncols, const int nrows,
|
|
1162
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1163
|
+
dpct::queue_ptr stream) {
|
|
1164
|
+
GGML_ASSERT(ncols % QK8_0 == 0);
|
|
1165
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
1166
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1167
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
1168
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1169
|
+
cgh.parallel_for(
|
|
1170
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1171
|
+
[=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1172
|
+
mul_mat_vec_q_ncols<QK8_0, QI8_0, block_q8_0,
|
|
1173
|
+
VDR_Q8_0_Q8_1_MMVQ, vec_dot_q8_0_q8_1, ncols_dst>(
|
|
1174
|
+
vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
|
|
1175
|
+
});
|
|
1176
|
+
});
|
|
1177
|
+
}
|
|
1178
|
+
|
|
1179
|
+
static void mul_mat_vec_q8_0_q8_1_sycl_switch_ncols(
|
|
1180
|
+
const void * vx, const void * vy, float * dst,
|
|
1181
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
1182
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1183
|
+
dpct::queue_ptr stream) {
|
|
1184
|
+
switch (ncols_dst) {
|
|
1185
|
+
case 1: mul_mat_vec_q8_0_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
1186
|
+
case 2: mul_mat_vec_q8_0_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1187
|
+
case 3: mul_mat_vec_q8_0_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1188
|
+
case 4: mul_mat_vec_q8_0_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1189
|
+
case 5: mul_mat_vec_q8_0_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1190
|
+
case 6: mul_mat_vec_q8_0_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1191
|
+
case 7: mul_mat_vec_q8_0_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1192
|
+
case 8: mul_mat_vec_q8_0_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1193
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q8_0 multi-col MMVQ", ncols_dst);
|
|
1194
|
+
}
|
|
1195
|
+
}
|
|
1196
|
+
|
|
1197
|
+
static void mul_mat_vec_q1_0_q8_1_sycl(const void * vx, const void * vy,
|
|
1198
|
+
float * dst, const int ncols,
|
|
1199
|
+
const int nrows,
|
|
1200
|
+
dpct::queue_ptr stream) {
|
|
1201
|
+
GGML_ASSERT(ncols % QK1_0 == 0);
|
|
1202
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
1203
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1204
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
1205
|
+
|
|
1206
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1207
|
+
cgh.parallel_for(
|
|
1208
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1209
|
+
[=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1210
|
+
mul_mat_vec_q<QK1_0, QI1_0, block_q1_0,
|
|
1211
|
+
VDR_Q1_0_Q8_1_MMVQ, vec_dot_q1_0_q8_1>(
|
|
1212
|
+
vx, vy, dst, ncols, nrows, item_ct1);
|
|
1213
|
+
});
|
|
1214
|
+
});
|
|
1215
|
+
}
|
|
1216
|
+
|
|
1217
|
+
template <int ncols_dst>
|
|
1218
|
+
static void mul_mat_vec_q1_0_q8_1_sycl_ncols(
|
|
1219
|
+
const void * vx, const void * vy, float * dst,
|
|
1220
|
+
const int ncols, const int nrows,
|
|
1221
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1222
|
+
dpct::queue_ptr stream) {
|
|
1223
|
+
GGML_ASSERT(ncols % QK1_0 == 0);
|
|
1224
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
1225
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1226
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
1227
|
+
|
|
1228
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1229
|
+
cgh.parallel_for(
|
|
1230
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1231
|
+
[=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1232
|
+
mul_mat_vec_q_ncols<QK1_0, QI1_0, block_q1_0,
|
|
1233
|
+
VDR_Q1_0_Q8_1_MMVQ, vec_dot_q1_0_q8_1, ncols_dst>(
|
|
1234
|
+
vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
|
|
1235
|
+
});
|
|
1236
|
+
});
|
|
1237
|
+
}
|
|
1238
|
+
|
|
1239
|
+
static void mul_mat_vec_q1_0_q8_1_sycl_switch_ncols(
|
|
1240
|
+
const void * vx, const void * vy, float * dst,
|
|
1241
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
1242
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1243
|
+
dpct::queue_ptr stream) {
|
|
1244
|
+
switch (ncols_dst) {
|
|
1245
|
+
case 1: mul_mat_vec_q1_0_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
1246
|
+
case 2: mul_mat_vec_q1_0_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1247
|
+
case 3: mul_mat_vec_q1_0_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1248
|
+
case 4: mul_mat_vec_q1_0_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1249
|
+
case 5: mul_mat_vec_q1_0_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1250
|
+
case 6: mul_mat_vec_q1_0_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1251
|
+
case 7: mul_mat_vec_q1_0_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1252
|
+
case 8: mul_mat_vec_q1_0_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1253
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q1_0 multi-col MMVQ", ncols_dst);
|
|
1254
|
+
}
|
|
1255
|
+
}
|
|
1256
|
+
|
|
689
1257
|
static void mul_mat_vec_q2_K_q8_1_sycl(const void *vx, const void *vy,
|
|
690
1258
|
float *dst, const int ncols,
|
|
691
1259
|
const int nrows,
|
|
@@ -710,6 +1278,45 @@ static void mul_mat_vec_q2_K_q8_1_sycl(const void *vx, const void *vy,
|
|
|
710
1278
|
}
|
|
711
1279
|
}
|
|
712
1280
|
|
|
1281
|
+
template <int ncols_dst>
|
|
1282
|
+
static void mul_mat_vec_q2_K_q8_1_sycl_ncols(
|
|
1283
|
+
const void * vx, const void * vy, float * dst,
|
|
1284
|
+
const int ncols, const int nrows,
|
|
1285
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1286
|
+
dpct::queue_ptr stream) {
|
|
1287
|
+
GGML_ASSERT(ncols % QK_K == 0);
|
|
1288
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
1289
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1290
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
1291
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1292
|
+
cgh.parallel_for(
|
|
1293
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1294
|
+
[=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1295
|
+
mul_mat_vec_q_ncols<QK_K, QI2_K, block_q2_K,
|
|
1296
|
+
VDR_Q2_K_Q8_1_MMVQ, vec_dot_q2_K_q8_1, ncols_dst>(
|
|
1297
|
+
vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
|
|
1298
|
+
});
|
|
1299
|
+
});
|
|
1300
|
+
}
|
|
1301
|
+
|
|
1302
|
+
static void mul_mat_vec_q2_K_q8_1_sycl_switch_ncols(
|
|
1303
|
+
const void * vx, const void * vy, float * dst,
|
|
1304
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
1305
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1306
|
+
dpct::queue_ptr stream) {
|
|
1307
|
+
switch (ncols_dst) {
|
|
1308
|
+
case 1: mul_mat_vec_q2_K_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
1309
|
+
case 2: mul_mat_vec_q2_K_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1310
|
+
case 3: mul_mat_vec_q2_K_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1311
|
+
case 4: mul_mat_vec_q2_K_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1312
|
+
case 5: mul_mat_vec_q2_K_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1313
|
+
case 6: mul_mat_vec_q2_K_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1314
|
+
case 7: mul_mat_vec_q2_K_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1315
|
+
case 8: mul_mat_vec_q2_K_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1316
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q2_K multi-col MMVQ", ncols_dst);
|
|
1317
|
+
}
|
|
1318
|
+
}
|
|
1319
|
+
|
|
713
1320
|
static void mul_mat_vec_q3_K_q8_1_sycl(const void *vx, const void *vy,
|
|
714
1321
|
float *dst, const int ncols,
|
|
715
1322
|
const int nrows,
|
|
@@ -734,6 +1341,104 @@ static void mul_mat_vec_q3_K_q8_1_sycl(const void *vx, const void *vy,
|
|
|
734
1341
|
}
|
|
735
1342
|
}
|
|
736
1343
|
|
|
1344
|
+
static void reorder_mul_mat_vec_q3_k_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols,
|
|
1345
|
+
const int nrows, dpct::queue_ptr stream) {
|
|
1346
|
+
GGML_ASSERT(ncols % QK_K == 0);
|
|
1347
|
+
|
|
1348
|
+
// Round up to a whole number of subgroup-sized workgroups; out-of-range rows are skipped inside the kernel.
|
|
1349
|
+
constexpr size_t num_subgroups = WARP_SIZE;
|
|
1350
|
+
const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
|
|
1351
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1352
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
1353
|
+
|
|
1354
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1355
|
+
cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1356
|
+
[=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1357
|
+
mul_mat_vec_q_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q3_K>>(vx, vy, dst, ncols, nrows,
|
|
1358
|
+
nd_item);
|
|
1359
|
+
});
|
|
1360
|
+
});
|
|
1361
|
+
}
|
|
1362
|
+
|
|
1363
|
+
template <int ncols_dst>
|
|
1364
|
+
static void reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols(
|
|
1365
|
+
const void * vx, const void * vy, float * dst,
|
|
1366
|
+
const int ncols, const int nrows,
|
|
1367
|
+
const int stride_col_y_bytes, const int stride_col_dst,
|
|
1368
|
+
dpct::queue_ptr stream) {
|
|
1369
|
+
GGML_ASSERT(ncols % QK_K == 0);
|
|
1370
|
+
constexpr size_t num_subgroups = WARP_SIZE;
|
|
1371
|
+
const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
|
|
1372
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1373
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
1374
|
+
|
|
1375
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1376
|
+
cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1377
|
+
[=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1378
|
+
mul_mat_vec_q_reorder_ncols<reorder_vec_dot_q_sycl<GGML_TYPE_Q3_K>, ncols_dst>(
|
|
1379
|
+
vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, nd_item);
|
|
1380
|
+
});
|
|
1381
|
+
});
|
|
1382
|
+
}
|
|
1383
|
+
|
|
1384
|
+
static void reorder_mul_mat_vec_q3_k_q8_1_sycl_switch_ncols(
|
|
1385
|
+
const void * vx, const void * vy, float * dst,
|
|
1386
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
1387
|
+
const int stride_col_y_bytes, const int stride_col_dst,
|
|
1388
|
+
dpct::queue_ptr stream) {
|
|
1389
|
+
switch (ncols_dst) {
|
|
1390
|
+
case 1: reorder_mul_mat_vec_q3_k_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
1391
|
+
case 2: reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1392
|
+
case 3: reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1393
|
+
case 4: reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1394
|
+
case 5: reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1395
|
+
case 6: reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1396
|
+
case 7: reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1397
|
+
case 8: reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1398
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q3_K reorder multi-col MMVQ", ncols_dst);
|
|
1399
|
+
}
|
|
1400
|
+
}
|
|
1401
|
+
|
|
1402
|
+
template <int ncols_dst>
|
|
1403
|
+
static void mul_mat_vec_q3_K_q8_1_sycl_ncols(
|
|
1404
|
+
const void * vx, const void * vy, float * dst,
|
|
1405
|
+
const int ncols, const int nrows,
|
|
1406
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1407
|
+
dpct::queue_ptr stream) {
|
|
1408
|
+
GGML_ASSERT(ncols % QK_K == 0);
|
|
1409
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
1410
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1411
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
1412
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1413
|
+
cgh.parallel_for(
|
|
1414
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1415
|
+
[=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1416
|
+
mul_mat_vec_q_ncols<QK_K, QI3_K, block_q3_K,
|
|
1417
|
+
VDR_Q3_K_Q8_1_MMVQ, vec_dot_q3_K_q8_1, ncols_dst>(
|
|
1418
|
+
vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
|
|
1419
|
+
});
|
|
1420
|
+
});
|
|
1421
|
+
}
|
|
1422
|
+
|
|
1423
|
+
static void mul_mat_vec_q3_K_q8_1_sycl_switch_ncols(
|
|
1424
|
+
const void * vx, const void * vy, float * dst,
|
|
1425
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
1426
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1427
|
+
dpct::queue_ptr stream) {
|
|
1428
|
+
switch (ncols_dst) {
|
|
1429
|
+
case 1: mul_mat_vec_q3_K_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
1430
|
+
case 2: mul_mat_vec_q3_K_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1431
|
+
case 3: mul_mat_vec_q3_K_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1432
|
+
case 4: mul_mat_vec_q3_K_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1433
|
+
case 5: mul_mat_vec_q3_K_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1434
|
+
case 6: mul_mat_vec_q3_K_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1435
|
+
case 7: mul_mat_vec_q3_K_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1436
|
+
case 8: mul_mat_vec_q3_K_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1437
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q3_K multi-col MMVQ", ncols_dst);
|
|
1438
|
+
}
|
|
1439
|
+
}
|
|
1440
|
+
|
|
1441
|
+
|
|
737
1442
|
static void mul_mat_vec_q4_K_q8_1_sycl(const void *vx, const void *vy,
|
|
738
1443
|
float *dst, const int ncols,
|
|
739
1444
|
const int nrows,
|
|
@@ -758,19 +1463,63 @@ static void mul_mat_vec_q4_K_q8_1_sycl(const void *vx, const void *vy,
|
|
|
758
1463
|
}
|
|
759
1464
|
}
|
|
760
1465
|
|
|
1466
|
+
template <int ncols_dst>
|
|
1467
|
+
static void mul_mat_vec_q4_K_q8_1_sycl_ncols(
|
|
1468
|
+
const void * vx, const void * vy, float * dst,
|
|
1469
|
+
const int ncols, const int nrows,
|
|
1470
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1471
|
+
dpct::queue_ptr stream) {
|
|
1472
|
+
GGML_ASSERT(ncols % QK_K == 0);
|
|
1473
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
1474
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1475
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
1476
|
+
|
|
1477
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1478
|
+
cgh.parallel_for(
|
|
1479
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1480
|
+
[=](sycl::nd_item<3> item_ct1)
|
|
1481
|
+
[[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1482
|
+
mul_mat_vec_q_ncols<QK_K, QI4_K, block_q4_K,
|
|
1483
|
+
VDR_Q4_K_Q8_1_MMVQ,
|
|
1484
|
+
vec_dot_q4_K_q8_1,
|
|
1485
|
+
ncols_dst>(
|
|
1486
|
+
vx, vy, dst, ncols, nrows,
|
|
1487
|
+
stride_col_y, stride_col_dst, item_ct1);
|
|
1488
|
+
});
|
|
1489
|
+
});
|
|
1490
|
+
}
|
|
1491
|
+
|
|
1492
|
+
static void mul_mat_vec_q4_K_q8_1_sycl_switch_ncols(
|
|
1493
|
+
const void * vx, const void * vy, float * dst,
|
|
1494
|
+
const int ncols, const int nrows,
|
|
1495
|
+
const int ncols_dst,
|
|
1496
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1497
|
+
dpct::queue_ptr stream) {
|
|
1498
|
+
switch (ncols_dst) {
|
|
1499
|
+
case 1: mul_mat_vec_q4_K_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
1500
|
+
case 2: mul_mat_vec_q4_K_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1501
|
+
case 3: mul_mat_vec_q4_K_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1502
|
+
case 4: mul_mat_vec_q4_K_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1503
|
+
case 5: mul_mat_vec_q4_K_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1504
|
+
case 6: mul_mat_vec_q4_K_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1505
|
+
case 7: mul_mat_vec_q4_K_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1506
|
+
case 8: mul_mat_vec_q4_K_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1507
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q4_K multi-col MMVQ", ncols_dst);
|
|
1508
|
+
}
|
|
1509
|
+
}
|
|
1510
|
+
|
|
761
1511
|
static void reorder_mul_mat_vec_q4_k_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols,
|
|
762
1512
|
const int nrows, dpct::queue_ptr stream) {
|
|
763
1513
|
GGML_ASSERT(ncols % QK_K == 0);
|
|
764
1514
|
|
|
765
|
-
|
|
766
|
-
constexpr size_t num_subgroups =
|
|
767
|
-
|
|
768
|
-
|
|
769
|
-
const sycl::range<3>
|
|
770
|
-
const sycl::range<3> workgroup_size(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
1515
|
+
// Round up to a whole number of subgroup-sized workgroups; out-of-range rows are skipped inside the kernel.
|
|
1516
|
+
constexpr size_t num_subgroups = WARP_SIZE;
|
|
1517
|
+
const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
|
|
1518
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1519
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
771
1520
|
|
|
772
1521
|
stream->submit([&](sycl::handler & cgh) {
|
|
773
|
-
cgh.parallel_for(sycl::nd_range<3>(
|
|
1522
|
+
cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
774
1523
|
[=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
775
1524
|
mul_mat_vec_q_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q4_K>>(vx, vy, dst, ncols,
|
|
776
1525
|
nrows, nd_item);
|
|
@@ -778,6 +1527,45 @@ static void reorder_mul_mat_vec_q4_k_q8_1_sycl(const void * vx, const void * vy,
|
|
|
778
1527
|
});
|
|
779
1528
|
}
|
|
780
1529
|
|
|
1530
|
+
template <int ncols_dst>
|
|
1531
|
+
static void reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols(
|
|
1532
|
+
const void * vx, const void * vy, float * dst,
|
|
1533
|
+
const int ncols, const int nrows,
|
|
1534
|
+
const int stride_col_y_bytes, const int stride_col_dst,
|
|
1535
|
+
dpct::queue_ptr stream) {
|
|
1536
|
+
GGML_ASSERT(ncols % QK_K == 0);
|
|
1537
|
+
|
|
1538
|
+
constexpr size_t num_subgroups = WARP_SIZE;
|
|
1539
|
+
const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
|
|
1540
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1541
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
1542
|
+
|
|
1543
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1544
|
+
cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1545
|
+
[=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1546
|
+
mul_mat_vec_q_reorder_ncols<reorder_vec_dot_q_sycl<GGML_TYPE_Q4_K>, ncols_dst>(
|
|
1547
|
+
vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, nd_item);
|
|
1548
|
+
});
|
|
1549
|
+
});
|
|
1550
|
+
}
|
|
1551
|
+
|
|
1552
|
+
static void reorder_mul_mat_vec_q4_k_q8_1_sycl_switch_ncols(
|
|
1553
|
+
const void * vx, const void * vy, float * dst,
|
|
1554
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
1555
|
+
const int stride_col_y_bytes, const int stride_col_dst,
|
|
1556
|
+
dpct::queue_ptr stream) {
|
|
1557
|
+
switch (ncols_dst) {
|
|
1558
|
+
case 1: reorder_mul_mat_vec_q4_k_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
1559
|
+
case 2: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1560
|
+
case 3: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1561
|
+
case 4: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1562
|
+
case 5: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1563
|
+
case 6: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1564
|
+
case 7: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1565
|
+
case 8: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1566
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q4_K reorder multi-col MMVQ", ncols_dst);
|
|
1567
|
+
}
|
|
1568
|
+
}
|
|
781
1569
|
|
|
782
1570
|
static void mul_mat_vec_q5_K_q8_1_sycl(const void *vx, const void *vy,
|
|
783
1571
|
float *dst, const int ncols,
|
|
@@ -803,24 +1591,167 @@ static void mul_mat_vec_q5_K_q8_1_sycl(const void *vx, const void *vy,
|
|
|
803
1591
|
}
|
|
804
1592
|
}
|
|
805
1593
|
|
|
1594
|
+
template <int ncols_dst>
|
|
1595
|
+
static void mul_mat_vec_q5_K_q8_1_sycl_ncols(
|
|
1596
|
+
const void * vx, const void * vy, float * dst,
|
|
1597
|
+
const int ncols, const int nrows,
|
|
1598
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1599
|
+
dpct::queue_ptr stream) {
|
|
1600
|
+
GGML_ASSERT(ncols % QK_K == 0);
|
|
1601
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
1602
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1603
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
1604
|
+
|
|
1605
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1606
|
+
cgh.parallel_for(
|
|
1607
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1608
|
+
[=](sycl::nd_item<3> item_ct1)
|
|
1609
|
+
[[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1610
|
+
mul_mat_vec_q_ncols<QK_K, QI5_K, block_q5_K,
|
|
1611
|
+
VDR_Q5_K_Q8_1_MMVQ,
|
|
1612
|
+
vec_dot_q5_K_q8_1,
|
|
1613
|
+
ncols_dst>(
|
|
1614
|
+
vx, vy, dst, ncols, nrows,
|
|
1615
|
+
stride_col_y, stride_col_dst, item_ct1);
|
|
1616
|
+
});
|
|
1617
|
+
});
|
|
1618
|
+
}
|
|
1619
|
+
|
|
1620
|
+
static void mul_mat_vec_q5_K_q8_1_sycl_switch_ncols(
|
|
1621
|
+
const void * vx, const void * vy, float * dst,
|
|
1622
|
+
const int ncols, const int nrows,
|
|
1623
|
+
const int ncols_dst,
|
|
1624
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1625
|
+
dpct::queue_ptr stream) {
|
|
1626
|
+
switch (ncols_dst) {
|
|
1627
|
+
case 1: mul_mat_vec_q5_K_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
1628
|
+
case 2: mul_mat_vec_q5_K_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1629
|
+
case 3: mul_mat_vec_q5_K_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1630
|
+
case 4: mul_mat_vec_q5_K_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1631
|
+
case 5: mul_mat_vec_q5_K_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1632
|
+
case 6: mul_mat_vec_q5_K_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1633
|
+
case 7: mul_mat_vec_q5_K_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1634
|
+
case 8: mul_mat_vec_q5_K_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1635
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q5_K multi-col MMVQ", ncols_dst);
|
|
1636
|
+
}
|
|
1637
|
+
}
|
|
1638
|
+
|
|
1639
|
+
static void reorder_mul_mat_vec_q5_k_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols,
|
|
1640
|
+
const int nrows, dpct::queue_ptr stream) {
|
|
1641
|
+
GGML_ASSERT(ncols % QK_K == 0);
|
|
1642
|
+
|
|
1643
|
+
constexpr size_t num_subgroups = WARP_SIZE;
|
|
1644
|
+
const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
|
|
1645
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1646
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
1647
|
+
|
|
1648
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1649
|
+
cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1650
|
+
[=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1651
|
+
mul_mat_vec_q_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q5_K>>(vx, vy, dst, ncols,
|
|
1652
|
+
nrows, nd_item);
|
|
1653
|
+
});
|
|
1654
|
+
});
|
|
1655
|
+
}
|
|
1656
|
+
|
|
1657
|
+
template <int ncols_dst>
|
|
1658
|
+
static void reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols(
|
|
1659
|
+
const void * vx, const void * vy, float * dst,
|
|
1660
|
+
const int ncols, const int nrows,
|
|
1661
|
+
const int stride_col_y_bytes, const int stride_col_dst,
|
|
1662
|
+
dpct::queue_ptr stream) {
|
|
1663
|
+
GGML_ASSERT(ncols % QK_K == 0);
|
|
1664
|
+
|
|
1665
|
+
constexpr size_t num_subgroups = WARP_SIZE;
|
|
1666
|
+
const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
|
|
1667
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1668
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
1669
|
+
|
|
1670
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1671
|
+
cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1672
|
+
[=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1673
|
+
mul_mat_vec_q_reorder_ncols<reorder_vec_dot_q_sycl<GGML_TYPE_Q5_K>, ncols_dst>(
|
|
1674
|
+
vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, nd_item);
|
|
1675
|
+
});
|
|
1676
|
+
});
|
|
1677
|
+
}
|
|
1678
|
+
|
|
1679
|
+
static void reorder_mul_mat_vec_q5_k_q8_1_sycl_switch_ncols(
|
|
1680
|
+
const void * vx, const void * vy, float * dst,
|
|
1681
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
1682
|
+
const int stride_col_y_bytes, const int stride_col_dst,
|
|
1683
|
+
dpct::queue_ptr stream) {
|
|
1684
|
+
switch (ncols_dst) {
|
|
1685
|
+
case 1: reorder_mul_mat_vec_q5_k_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
1686
|
+
case 2: reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1687
|
+
case 3: reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1688
|
+
case 4: reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1689
|
+
case 5: reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1690
|
+
case 6: reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1691
|
+
case 7: reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1692
|
+
case 8: reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1693
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q5_K reorder multi-col MMVQ", ncols_dst);
|
|
1694
|
+
}
|
|
1695
|
+
}
|
|
1696
|
+
|
|
806
1697
|
static void reorder_mul_mat_vec_q6_k_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols,
|
|
807
1698
|
const int nrows, dpct::queue_ptr stream) {
|
|
808
1699
|
GGML_ASSERT(ncols % QK_K == 0);
|
|
809
|
-
|
|
810
|
-
constexpr size_t num_subgroups =
|
|
811
|
-
|
|
1700
|
+
// Round up to a whole number of subgroup-sized workgroups; out-of-range rows are skipped inside the kernel.
|
|
1701
|
+
constexpr size_t num_subgroups = WARP_SIZE;
|
|
1702
|
+
const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
|
|
1703
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1704
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
812
1705
|
|
|
813
|
-
const sycl::range<3> global_size(1, GGML_SYCL_MMV_Y, block_num_y * WARP_SIZE);
|
|
814
|
-
const sycl::range<3> workgroup_size(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
815
1706
|
|
|
816
1707
|
stream->submit([&](sycl::handler & cgh) {
|
|
817
|
-
cgh.parallel_for(sycl::nd_range<3>(
|
|
1708
|
+
cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
818
1709
|
[=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
819
1710
|
mul_mat_vec_q_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q6_K>>(vx, vy, dst, ncols, nrows,
|
|
820
1711
|
nd_item);
|
|
821
1712
|
});
|
|
822
1713
|
});
|
|
823
1714
|
}
|
|
1715
|
+
|
|
1716
|
+
template <int ncols_dst>
|
|
1717
|
+
static void reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols(
|
|
1718
|
+
const void * vx, const void * vy, float * dst,
|
|
1719
|
+
const int ncols, const int nrows,
|
|
1720
|
+
const int stride_col_y_bytes, const int stride_col_dst,
|
|
1721
|
+
dpct::queue_ptr stream) {
|
|
1722
|
+
GGML_ASSERT(ncols % QK_K == 0);
|
|
1723
|
+
constexpr size_t num_subgroups = WARP_SIZE;
|
|
1724
|
+
const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
|
|
1725
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1726
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
|
1727
|
+
|
|
1728
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1729
|
+
cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1730
|
+
[=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1731
|
+
mul_mat_vec_q_reorder_ncols<reorder_vec_dot_q_sycl<GGML_TYPE_Q6_K>, ncols_dst>(
|
|
1732
|
+
vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, nd_item);
|
|
1733
|
+
});
|
|
1734
|
+
});
|
|
1735
|
+
}
|
|
1736
|
+
|
|
1737
|
+
static void reorder_mul_mat_vec_q6_k_q8_1_sycl_switch_ncols(
|
|
1738
|
+
const void * vx, const void * vy, float * dst,
|
|
1739
|
+
const int ncols, const int nrows, const int ncols_dst,
|
|
1740
|
+
const int stride_col_y_bytes, const int stride_col_dst,
|
|
1741
|
+
dpct::queue_ptr stream) {
|
|
1742
|
+
switch (ncols_dst) {
|
|
1743
|
+
case 1: reorder_mul_mat_vec_q6_k_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
1744
|
+
case 2: reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1745
|
+
case 3: reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1746
|
+
case 4: reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1747
|
+
case 5: reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1748
|
+
case 6: reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1749
|
+
case 7: reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1750
|
+
case 8: reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
|
|
1751
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q6_K reorder multi-col MMVQ", ncols_dst);
|
|
1752
|
+
}
|
|
1753
|
+
}
|
|
1754
|
+
|
|
824
1755
|
static void mul_mat_vec_q6_K_q8_1_sycl(const void *vx, const void *vy,
|
|
825
1756
|
float *dst, const int ncols,
|
|
826
1757
|
const int nrows,
|
|
@@ -845,6 +1776,51 @@ static void mul_mat_vec_q6_K_q8_1_sycl(const void *vx, const void *vy,
|
|
|
845
1776
|
}
|
|
846
1777
|
}
|
|
847
1778
|
|
|
1779
|
+
template <int ncols_dst>
|
|
1780
|
+
static void mul_mat_vec_q6_K_q8_1_sycl_ncols(
|
|
1781
|
+
const void * vx, const void * vy, float * dst,
|
|
1782
|
+
const int ncols, const int nrows,
|
|
1783
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1784
|
+
dpct::queue_ptr stream) {
|
|
1785
|
+
GGML_ASSERT(ncols % QK_K == 0);
|
|
1786
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
1787
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
1788
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
1789
|
+
|
|
1790
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
1791
|
+
cgh.parallel_for(
|
|
1792
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
1793
|
+
[=](sycl::nd_item<3> item_ct1)
|
|
1794
|
+
[[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
1795
|
+
mul_mat_vec_q_ncols<QK_K, QI6_K, block_q6_K,
|
|
1796
|
+
VDR_Q6_K_Q8_1_MMVQ,
|
|
1797
|
+
vec_dot_q6_K_q8_1,
|
|
1798
|
+
ncols_dst>(
|
|
1799
|
+
vx, vy, dst, ncols, nrows,
|
|
1800
|
+
stride_col_y, stride_col_dst, item_ct1);
|
|
1801
|
+
});
|
|
1802
|
+
});
|
|
1803
|
+
}
|
|
1804
|
+
|
|
1805
|
+
static void mul_mat_vec_q6_K_q8_1_sycl_switch_ncols(
|
|
1806
|
+
const void * vx, const void * vy, float * dst,
|
|
1807
|
+
const int ncols, const int nrows,
|
|
1808
|
+
const int ncols_dst,
|
|
1809
|
+
const int stride_col_y, const int stride_col_dst,
|
|
1810
|
+
dpct::queue_ptr stream) {
|
|
1811
|
+
switch (ncols_dst) {
|
|
1812
|
+
case 1: mul_mat_vec_q6_K_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
1813
|
+
case 2: mul_mat_vec_q6_K_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1814
|
+
case 3: mul_mat_vec_q6_K_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1815
|
+
case 4: mul_mat_vec_q6_K_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1816
|
+
case 5: mul_mat_vec_q6_K_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1817
|
+
case 6: mul_mat_vec_q6_K_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1818
|
+
case 7: mul_mat_vec_q6_K_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1819
|
+
case 8: mul_mat_vec_q6_K_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
1820
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for Q6_K multi-col MMVQ", ncols_dst);
|
|
1821
|
+
}
|
|
1822
|
+
}
|
|
1823
|
+
|
|
848
1824
|
|
|
849
1825
|
static void mul_mat_vec_iq2_xxs_q8_1_sycl(const void *vx, const void *vy,
|
|
850
1826
|
float *dst, const int ncols,
|
|
@@ -1041,6 +2017,51 @@ static void mul_mat_vec_iq4_xs_q8_1_sycl(const void *vx, const void *vy,
|
|
|
1041
2017
|
}
|
|
1042
2018
|
}
|
|
1043
2019
|
|
|
2020
|
+
template <int ncols_dst>
|
|
2021
|
+
static void mul_mat_vec_iq4_xs_q8_1_sycl_ncols(
|
|
2022
|
+
const void * vx, const void * vy, float * dst,
|
|
2023
|
+
const int ncols, const int nrows,
|
|
2024
|
+
const int stride_col_y, const int stride_col_dst,
|
|
2025
|
+
dpct::queue_ptr stream) {
|
|
2026
|
+
GGML_ASSERT(ncols % QK_K == 0);
|
|
2027
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
2028
|
+
const sycl::range<3> block_nums(1, 1, block_num_y);
|
|
2029
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
2030
|
+
|
|
2031
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
2032
|
+
cgh.parallel_for(
|
|
2033
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
2034
|
+
[=](sycl::nd_item<3> item_ct1)
|
|
2035
|
+
[[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
2036
|
+
mul_mat_vec_q_ncols<QK_K, QI4_XS/4, block_iq4_xs,
|
|
2037
|
+
1,
|
|
2038
|
+
vec_dot_iq4_xs_q8_1,
|
|
2039
|
+
ncols_dst>(
|
|
2040
|
+
vx, vy, dst, ncols, nrows,
|
|
2041
|
+
stride_col_y, stride_col_dst, item_ct1);
|
|
2042
|
+
});
|
|
2043
|
+
});
|
|
2044
|
+
}
|
|
2045
|
+
|
|
2046
|
+
static void mul_mat_vec_iq4_xs_q8_1_sycl_switch_ncols(
|
|
2047
|
+
const void * vx, const void * vy, float * dst,
|
|
2048
|
+
const int ncols, const int nrows,
|
|
2049
|
+
const int ncols_dst,
|
|
2050
|
+
const int stride_col_y, const int stride_col_dst,
|
|
2051
|
+
dpct::queue_ptr stream) {
|
|
2052
|
+
switch (ncols_dst) {
|
|
2053
|
+
case 1: mul_mat_vec_iq4_xs_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
|
|
2054
|
+
case 2: mul_mat_vec_iq4_xs_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
2055
|
+
case 3: mul_mat_vec_iq4_xs_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
2056
|
+
case 4: mul_mat_vec_iq4_xs_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
2057
|
+
case 5: mul_mat_vec_iq4_xs_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
2058
|
+
case 6: mul_mat_vec_iq4_xs_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
2059
|
+
case 7: mul_mat_vec_iq4_xs_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
2060
|
+
case 8: mul_mat_vec_iq4_xs_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
|
|
2061
|
+
default: GGML_ABORT("unsupported ncols_dst=%d for IQ4_XS multi-col MMVQ", ncols_dst);
|
|
2062
|
+
}
|
|
2063
|
+
}
|
|
2064
|
+
|
|
1044
2065
|
void ggml_sycl_op_mul_mat_vec_q(ggml_backend_sycl_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1,
|
|
1045
2066
|
ggml_tensor * dst, const char * src0_dd_i, const float * src1_ddf_i,
|
|
1046
2067
|
const char * src1_ddq_i, float * dst_dd_i, const int64_t row_low,
|
|
@@ -1067,50 +2088,233 @@ void ggml_sycl_op_mul_mat_vec_q(ggml_backend_sycl_context & ctx, const ggml_tens
|
|
|
1067
2088
|
case GGML_TYPE_Q4_0:
|
|
1068
2089
|
if ((ggml_tensor_extra_gpu *) dst->src[0]->extra &&
|
|
1069
2090
|
((ggml_tensor_extra_gpu *) dst->src[0]->extra)->optimized_feature.reorder) {
|
|
1070
|
-
|
|
1071
|
-
|
|
1072
|
-
|
|
2091
|
+
if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2092
|
+
const int stride_col_y_bytes = src1_padded_col_size * q8_1_ts / q8_1_bs;
|
|
2093
|
+
const int stride_col_dst = dst->ne[0];
|
|
2094
|
+
GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q4_0_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2095
|
+
reorder_mul_mat_vec_q4_0_q8_1_sycl_switch_ncols(
|
|
2096
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2097
|
+
src1_ncols, stride_col_y_bytes, stride_col_dst, stream);
|
|
2098
|
+
return;
|
|
2099
|
+
} else {
|
|
2100
|
+
GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q4_0_q8_1_sycl\n");
|
|
2101
|
+
reorder_mul_mat_vec_q4_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2102
|
+
}
|
|
2103
|
+
} else if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2104
|
+
const int stride_col_y = src1_padded_col_size / QK8_1;
|
|
2105
|
+
const int stride_col_dst = dst->ne[0];
|
|
2106
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q4_0_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2107
|
+
mul_mat_vec_q4_0_q8_1_sycl_switch_ncols(
|
|
2108
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2109
|
+
src1_ncols, stride_col_y, stride_col_dst, stream);
|
|
2110
|
+
return;
|
|
2111
|
+
} else if (i == 0 || src1_ncols == 1) {
|
|
1073
2112
|
GGML_SYCL_DEBUG("Calling mul_mat_vec_q4_0_q8_1_sycl\n");
|
|
1074
2113
|
mul_mat_vec_q4_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
1075
2114
|
}
|
|
1076
2115
|
break;
|
|
1077
2116
|
case GGML_TYPE_Q4_1:
|
|
1078
|
-
|
|
2117
|
+
if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2118
|
+
const int stride_col_y = src1_padded_col_size / QK8_1;
|
|
2119
|
+
const int stride_col_dst = dst->ne[0];
|
|
2120
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q4_1_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2121
|
+
mul_mat_vec_q4_1_q8_1_sycl_switch_ncols(
|
|
2122
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2123
|
+
src1_ncols, stride_col_y, stride_col_dst, stream);
|
|
2124
|
+
return;
|
|
2125
|
+
} else if (i == 0 || src1_ncols == 1) {
|
|
2126
|
+
mul_mat_vec_q4_1_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2127
|
+
}
|
|
1079
2128
|
break;
|
|
1080
2129
|
case GGML_TYPE_Q5_0:
|
|
1081
|
-
|
|
2130
|
+
if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2131
|
+
const int stride_col_y = src1_padded_col_size / QK8_1;
|
|
2132
|
+
const int stride_col_dst = dst->ne[0];
|
|
2133
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q5_0_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2134
|
+
mul_mat_vec_q5_0_q8_1_sycl_switch_ncols(
|
|
2135
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2136
|
+
src1_ncols, stride_col_y, stride_col_dst, stream);
|
|
2137
|
+
return;
|
|
2138
|
+
} else if (i == 0 || src1_ncols == 1) {
|
|
2139
|
+
mul_mat_vec_q5_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2140
|
+
}
|
|
1082
2141
|
break;
|
|
1083
2142
|
case GGML_TYPE_Q5_1:
|
|
1084
|
-
|
|
2143
|
+
if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2144
|
+
const int stride_col_y = src1_padded_col_size / QK8_1;
|
|
2145
|
+
const int stride_col_dst = dst->ne[0];
|
|
2146
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q5_1_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2147
|
+
mul_mat_vec_q5_1_q8_1_sycl_switch_ncols(
|
|
2148
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2149
|
+
src1_ncols, stride_col_y, stride_col_dst, stream);
|
|
2150
|
+
return;
|
|
2151
|
+
} else if (i == 0 || src1_ncols == 1) {
|
|
2152
|
+
mul_mat_vec_q5_1_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2153
|
+
}
|
|
1085
2154
|
break;
|
|
1086
2155
|
case GGML_TYPE_Q8_0:
|
|
1087
|
-
|
|
2156
|
+
if ((ggml_tensor_extra_gpu *) dst->src[0]->extra &&
|
|
2157
|
+
((ggml_tensor_extra_gpu *) dst->src[0]->extra)->optimized_feature.reorder) {
|
|
2158
|
+
if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2159
|
+
const int stride_col_y_bytes = src1_padded_col_size * q8_1_ts / q8_1_bs;
|
|
2160
|
+
const int stride_col_dst = dst->ne[0];
|
|
2161
|
+
GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q8_0_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2162
|
+
reorder_mul_mat_vec_q8_0_q8_1_sycl_switch_ncols(
|
|
2163
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2164
|
+
src1_ncols, stride_col_y_bytes, stride_col_dst, stream);
|
|
2165
|
+
return;
|
|
2166
|
+
} else {
|
|
2167
|
+
GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q8_0_q8_1_sycl\n");
|
|
2168
|
+
reorder_mul_mat_vec_q8_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2169
|
+
}
|
|
2170
|
+
} else if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2171
|
+
const int stride_col_y = src1_padded_col_size / QK8_1;
|
|
2172
|
+
const int stride_col_dst = dst->ne[0];
|
|
2173
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q8_0_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2174
|
+
mul_mat_vec_q8_0_q8_1_sycl_switch_ncols(
|
|
2175
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2176
|
+
src1_ncols, stride_col_y, stride_col_dst, stream);
|
|
2177
|
+
return;
|
|
2178
|
+
} else if (i == 0 || src1_ncols == 1) {
|
|
2179
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q8_0_q8_1_sycl\n");
|
|
2180
|
+
mul_mat_vec_q8_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2181
|
+
}
|
|
2182
|
+
break;
|
|
2183
|
+
case GGML_TYPE_Q1_0:
|
|
2184
|
+
if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2185
|
+
const int stride_col_y = src1_padded_col_size / QK8_1;
|
|
2186
|
+
const int stride_col_dst = dst->ne[0];
|
|
2187
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q1_0_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2188
|
+
mul_mat_vec_q1_0_q8_1_sycl_switch_ncols(
|
|
2189
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2190
|
+
src1_ncols, stride_col_y, stride_col_dst, stream);
|
|
2191
|
+
return;
|
|
2192
|
+
} else if (i == 0 || src1_ncols == 1) {
|
|
2193
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q1_0_q8_1_sycl\n");
|
|
2194
|
+
mul_mat_vec_q1_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2195
|
+
}
|
|
1088
2196
|
break;
|
|
1089
2197
|
case GGML_TYPE_Q2_K:
|
|
1090
|
-
|
|
2198
|
+
if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2199
|
+
const int stride_col_y = src1_padded_col_size / QK8_1;
|
|
2200
|
+
const int stride_col_dst = dst->ne[0];
|
|
2201
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q2_K_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2202
|
+
mul_mat_vec_q2_K_q8_1_sycl_switch_ncols(
|
|
2203
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2204
|
+
src1_ncols, stride_col_y, stride_col_dst, stream);
|
|
2205
|
+
return;
|
|
2206
|
+
} else if (i == 0 || src1_ncols == 1) {
|
|
2207
|
+
mul_mat_vec_q2_K_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2208
|
+
}
|
|
1091
2209
|
break;
|
|
1092
2210
|
case GGML_TYPE_Q3_K:
|
|
1093
|
-
|
|
2211
|
+
if ((ggml_tensor_extra_gpu *) dst->src[0]->extra &&
|
|
2212
|
+
((ggml_tensor_extra_gpu *) dst->src[0]->extra)->optimized_feature.reorder) {
|
|
2213
|
+
if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2214
|
+
const int stride_col_y_bytes = src1_padded_col_size * q8_1_ts / q8_1_bs;
|
|
2215
|
+
const int stride_col_dst = dst->ne[0];
|
|
2216
|
+
GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q3_k_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2217
|
+
reorder_mul_mat_vec_q3_k_q8_1_sycl_switch_ncols(
|
|
2218
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2219
|
+
src1_ncols, stride_col_y_bytes, stride_col_dst, stream);
|
|
2220
|
+
return;
|
|
2221
|
+
} else {
|
|
2222
|
+
GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q3_k_q8_1_sycl\n");
|
|
2223
|
+
reorder_mul_mat_vec_q3_k_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2224
|
+
}
|
|
2225
|
+
} else if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2226
|
+
const int stride_col_y = src1_padded_col_size / QK8_1;
|
|
2227
|
+
const int stride_col_dst = dst->ne[0];
|
|
2228
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q3_K_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2229
|
+
mul_mat_vec_q3_K_q8_1_sycl_switch_ncols(
|
|
2230
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2231
|
+
src1_ncols, stride_col_y, stride_col_dst, stream);
|
|
2232
|
+
return;
|
|
2233
|
+
} else if (i == 0 || src1_ncols == 1) {
|
|
2234
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q3_K_q8_1_sycl\n");
|
|
2235
|
+
mul_mat_vec_q3_K_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2236
|
+
}
|
|
1094
2237
|
break;
|
|
1095
2238
|
case GGML_TYPE_Q4_K:
|
|
1096
2239
|
if ((ggml_tensor_extra_gpu *) dst->src[0]->extra &&
|
|
1097
2240
|
((ggml_tensor_extra_gpu *) dst->src[0]->extra)->optimized_feature.reorder) {
|
|
1098
|
-
|
|
1099
|
-
|
|
1100
|
-
|
|
2241
|
+
if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2242
|
+
const int stride_col_y_bytes = src1_padded_col_size * q8_1_ts / q8_1_bs;
|
|
2243
|
+
const int stride_col_dst = dst->ne[0];
|
|
2244
|
+
GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q4_k_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2245
|
+
reorder_mul_mat_vec_q4_k_q8_1_sycl_switch_ncols(
|
|
2246
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2247
|
+
src1_ncols, stride_col_y_bytes, stride_col_dst, stream);
|
|
2248
|
+
return;
|
|
2249
|
+
} else {
|
|
2250
|
+
GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q4_k_q8_1_sycl\n");
|
|
2251
|
+
reorder_mul_mat_vec_q4_k_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2252
|
+
}
|
|
2253
|
+
} else if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2254
|
+
const int stride_col_y = src1_padded_col_size / QK8_1;
|
|
2255
|
+
const int stride_col_dst = dst->ne[0];
|
|
2256
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q4_K_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2257
|
+
mul_mat_vec_q4_K_q8_1_sycl_switch_ncols(
|
|
2258
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2259
|
+
src1_ncols, stride_col_y, stride_col_dst, stream);
|
|
2260
|
+
return;
|
|
2261
|
+
} else if (i == 0 || src1_ncols == 1) {
|
|
1101
2262
|
GGML_SYCL_DEBUG("Calling mul_mat_vec_q4_K_q8_1_sycl\n");
|
|
1102
2263
|
mul_mat_vec_q4_K_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
1103
2264
|
}
|
|
1104
2265
|
break;
|
|
1105
2266
|
case GGML_TYPE_Q5_K:
|
|
1106
|
-
|
|
2267
|
+
if ((ggml_tensor_extra_gpu *) dst->src[0]->extra &&
|
|
2268
|
+
((ggml_tensor_extra_gpu *) dst->src[0]->extra)->optimized_feature.reorder) {
|
|
2269
|
+
if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2270
|
+
const int stride_col_y_bytes = src1_padded_col_size * q8_1_ts / q8_1_bs;
|
|
2271
|
+
const int stride_col_dst = dst->ne[0];
|
|
2272
|
+
GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q5_k_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2273
|
+
reorder_mul_mat_vec_q5_k_q8_1_sycl_switch_ncols(
|
|
2274
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2275
|
+
src1_ncols, stride_col_y_bytes, stride_col_dst, stream);
|
|
2276
|
+
return;
|
|
2277
|
+
} else {
|
|
2278
|
+
GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q5_k_q8_1_sycl\n");
|
|
2279
|
+
reorder_mul_mat_vec_q5_k_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2280
|
+
}
|
|
2281
|
+
} else if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2282
|
+
const int stride_col_y = src1_padded_col_size / QK8_1;
|
|
2283
|
+
const int stride_col_dst = dst->ne[0];
|
|
2284
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q5_K_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2285
|
+
mul_mat_vec_q5_K_q8_1_sycl_switch_ncols(
|
|
2286
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2287
|
+
src1_ncols, stride_col_y, stride_col_dst, stream);
|
|
2288
|
+
return;
|
|
2289
|
+
} else if (i == 0 || src1_ncols == 1) {
|
|
2290
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q5_K_q8_1_sycl\n");
|
|
2291
|
+
mul_mat_vec_q5_K_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2292
|
+
}
|
|
1107
2293
|
break;
|
|
1108
2294
|
case GGML_TYPE_Q6_K:
|
|
1109
2295
|
if ((ggml_tensor_extra_gpu *) dst->src[0]->extra &&
|
|
1110
2296
|
((ggml_tensor_extra_gpu *) dst->src[0]->extra)->optimized_feature.reorder) {
|
|
1111
|
-
|
|
1112
|
-
|
|
1113
|
-
|
|
2297
|
+
if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2298
|
+
const int stride_col_y_bytes = src1_padded_col_size * q8_1_ts / q8_1_bs;
|
|
2299
|
+
const int stride_col_dst = dst->ne[0];
|
|
2300
|
+
GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q6_k_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2301
|
+
reorder_mul_mat_vec_q6_k_q8_1_sycl_switch_ncols(
|
|
2302
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2303
|
+
src1_ncols, stride_col_y_bytes, stride_col_dst, stream);
|
|
2304
|
+
return;
|
|
2305
|
+
} else {
|
|
2306
|
+
GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q6_k_q8_1_sycl\n");
|
|
2307
|
+
reorder_mul_mat_vec_q6_k_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2308
|
+
}
|
|
2309
|
+
} else if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2310
|
+
const int stride_col_y = src1_padded_col_size / QK8_1;
|
|
2311
|
+
const int stride_col_dst = dst->ne[0];
|
|
2312
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_q6_K_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2313
|
+
mul_mat_vec_q6_K_q8_1_sycl_switch_ncols(
|
|
2314
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2315
|
+
src1_ncols, stride_col_y, stride_col_dst, stream);
|
|
2316
|
+
return;
|
|
2317
|
+
} else if (i == 0 || src1_ncols == 1) {
|
|
1114
2318
|
GGML_SYCL_DEBUG("Calling mul_mat_vec_q6_k_q8_1_sycl\n");
|
|
1115
2319
|
mul_mat_vec_q6_K_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
1116
2320
|
}
|
|
@@ -1140,13 +2344,46 @@ void ggml_sycl_op_mul_mat_vec_q(ggml_backend_sycl_context & ctx, const ggml_tens
|
|
|
1140
2344
|
mul_mat_vec_iq4_nl_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
1141
2345
|
break;
|
|
1142
2346
|
case GGML_TYPE_IQ4_XS:
|
|
1143
|
-
|
|
2347
|
+
if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2348
|
+
const int stride_col_y = src1_padded_col_size / QK8_1;
|
|
2349
|
+
const int stride_col_dst = dst->ne[0];
|
|
2350
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_iq4_xs_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2351
|
+
mul_mat_vec_iq4_xs_q8_1_sycl_switch_ncols(
|
|
2352
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2353
|
+
src1_ncols, stride_col_y, stride_col_dst, stream);
|
|
2354
|
+
return;
|
|
2355
|
+
} else if (i == 0 || src1_ncols == 1) {
|
|
2356
|
+
mul_mat_vec_iq4_xs_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2357
|
+
}
|
|
1144
2358
|
break;
|
|
1145
2359
|
case GGML_TYPE_MXFP4:
|
|
1146
|
-
|
|
2360
|
+
if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2361
|
+
const int stride_col_y = src1_padded_col_size / QK8_1;
|
|
2362
|
+
const int stride_col_dst = dst->ne[0];
|
|
2363
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_mxfp4_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2364
|
+
mul_mat_vec_mxfp4_q8_1_sycl_switch_ncols(
|
|
2365
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2366
|
+
src1_ncols, stride_col_y, stride_col_dst, stream);
|
|
2367
|
+
return;
|
|
2368
|
+
} else if (i == 0 || src1_ncols == 1) {
|
|
2369
|
+
mul_mat_vec_mxfp4_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2370
|
+
}
|
|
2371
|
+
break;
|
|
2372
|
+
case GGML_TYPE_NVFP4:
|
|
2373
|
+
if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
|
|
2374
|
+
const int stride_col_y = src1_padded_col_size / QK8_1;
|
|
2375
|
+
const int stride_col_dst = dst->ne[0];
|
|
2376
|
+
GGML_SYCL_DEBUG("Calling mul_mat_vec_nvfp4_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
|
|
2377
|
+
mul_mat_vec_nvfp4_q8_1_sycl_switch_ncols(
|
|
2378
|
+
src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
|
|
2379
|
+
src1_ncols, stride_col_y, stride_col_dst, stream);
|
|
2380
|
+
return;
|
|
2381
|
+
} else if (i == 0 || src1_ncols == 1) {
|
|
2382
|
+
mul_mat_vec_nvfp4_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
|
|
2383
|
+
}
|
|
1147
2384
|
break;
|
|
1148
2385
|
default:
|
|
1149
|
-
GGML_ABORT("fatal error");
|
|
2386
|
+
GGML_ABORT("fatal error: unsupport data type=%s\n", ggml_type_name(src0->type));
|
|
1150
2387
|
}
|
|
1151
2388
|
}
|
|
1152
2389
|
GGML_UNUSED(src1);
|
|
@@ -1154,3 +2391,269 @@ void ggml_sycl_op_mul_mat_vec_q(ggml_backend_sycl_context & ctx, const ggml_tens
|
|
|
1154
2391
|
GGML_UNUSED(src1_ddf_i);
|
|
1155
2392
|
GGML_UNUSED(ctx);
|
|
1156
2393
|
}
|
|
2394
|
+
|
|
2395
|
+
// src1_row_stride: 0 for shared src1 (gate/up proj), else per-expert stride (down proj).
|
|
2396
|
+
template <int qk, int qi, typename block_q_t, int vdr, vec_dot_q_sycl_t vec_dot_q_sycl>
|
|
2397
|
+
static void mul_mat_vec_q_moe(
|
|
2398
|
+
const void * __restrict__ vx_base, const void * __restrict__ vy_base,
|
|
2399
|
+
float * __restrict__ dst_base, const int32_t * __restrict__ ids_dev,
|
|
2400
|
+
const int ncols, const int nrows,
|
|
2401
|
+
const size_t expert_weight_stride, const size_t dst_row_stride,
|
|
2402
|
+
const size_t src1_row_stride,
|
|
2403
|
+
const sycl::nd_item<3> & item_ct1) {
|
|
2404
|
+
|
|
2405
|
+
const int expert_idx = item_ct1.get_group(1);
|
|
2406
|
+
const int i02 = ids_dev[expert_idx];
|
|
2407
|
+
|
|
2408
|
+
const char * vx = (const char *) vx_base + (size_t) i02 * expert_weight_stride;
|
|
2409
|
+
const char * vy = (const char *) vy_base + (size_t) expert_idx * src1_row_stride;
|
|
2410
|
+
float * dst = (float *) ((char *) dst_base + (size_t) expert_idx * dst_row_stride);
|
|
2411
|
+
|
|
2412
|
+
const int row = item_ct1.get_group(2) * item_ct1.get_local_range(1) + item_ct1.get_local_id(1);
|
|
2413
|
+
|
|
2414
|
+
if (row >= nrows) {
|
|
2415
|
+
return;
|
|
2416
|
+
}
|
|
2417
|
+
|
|
2418
|
+
const int blocks_per_row = ncols / qk;
|
|
2419
|
+
constexpr int blocks_per_warp = (vdr * WARP_SIZE + qi - 1) / qi;
|
|
2420
|
+
|
|
2421
|
+
float tmp = 0.0f;
|
|
2422
|
+
|
|
2423
|
+
const block_q_t * x = (const block_q_t *) vx;
|
|
2424
|
+
const block_q8_1 * y = (const block_q8_1 *) vy;
|
|
2425
|
+
|
|
2426
|
+
for (int i = item_ct1.get_local_id(2) / (qi / vdr); i < blocks_per_row; i += blocks_per_warp) {
|
|
2427
|
+
const int ibx = row * blocks_per_row + i;
|
|
2428
|
+
const int iby = i * (qk / QK8_1);
|
|
2429
|
+
|
|
2430
|
+
for (size_t elem = 0; elem < qi / vdr; elem += WARP_SIZE) {
|
|
2431
|
+
const int iqs = elem + vdr * (item_ct1.get_local_id(2) % (qi / vdr));
|
|
2432
|
+
tmp += vec_dot_q_sycl(&x[ibx], &y[iby], iqs);
|
|
2433
|
+
}
|
|
2434
|
+
}
|
|
2435
|
+
|
|
2436
|
+
#pragma unroll
|
|
2437
|
+
for (int mask = WARP_SIZE / 2; mask > 0; mask >>= 1) {
|
|
2438
|
+
tmp += dpct::permute_sub_group_by_xor(item_ct1.get_sub_group(), tmp, mask);
|
|
2439
|
+
}
|
|
2440
|
+
|
|
2441
|
+
if (item_ct1.get_local_id(2) == 0) {
|
|
2442
|
+
dst[row] = tmp;
|
|
2443
|
+
}
|
|
2444
|
+
}
|
|
2445
|
+
|
|
2446
|
+
template <int qk, int qi, typename block_q_t, int vdr, vec_dot_q_sycl_t vec_dot_q_sycl>
|
|
2447
|
+
static void launch_mul_mat_vec_q_moe(
|
|
2448
|
+
const void * vx_base, const void * vy, const int32_t * ids_dev,
|
|
2449
|
+
float * dst_base, const int ncols, const int nrows, const int n_experts_used,
|
|
2450
|
+
const size_t expert_weight_stride, const size_t dst_row_stride,
|
|
2451
|
+
const size_t src1_row_stride,
|
|
2452
|
+
dpct::queue_ptr stream) {
|
|
2453
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
2454
|
+
const sycl::range<3> block_nums(1, (unsigned) n_experts_used, (unsigned) block_num_y);
|
|
2455
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
2456
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
2457
|
+
cgh.parallel_for(
|
|
2458
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
2459
|
+
[=](sycl::nd_item<3> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
2460
|
+
mul_mat_vec_q_moe<qk, qi, block_q_t, vdr, vec_dot_q_sycl>(
|
|
2461
|
+
vx_base, vy, dst_base, ids_dev, ncols, nrows,
|
|
2462
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, item);
|
|
2463
|
+
});
|
|
2464
|
+
});
|
|
2465
|
+
}
|
|
2466
|
+
|
|
2467
|
+
bool ggml_sycl_mul_mat_vec_q_id(
|
|
2468
|
+
enum ggml_type src0_type,
|
|
2469
|
+
const void * vx_base,
|
|
2470
|
+
const void * vy,
|
|
2471
|
+
const int32_t * ids_dev,
|
|
2472
|
+
float * dst_base,
|
|
2473
|
+
int ncols,
|
|
2474
|
+
int nrows,
|
|
2475
|
+
int n_experts_used,
|
|
2476
|
+
size_t expert_weight_stride,
|
|
2477
|
+
size_t dst_row_stride,
|
|
2478
|
+
size_t src1_row_stride,
|
|
2479
|
+
dpct::queue_ptr stream) {
|
|
2480
|
+
switch (src0_type) {
|
|
2481
|
+
case GGML_TYPE_Q4_0:
|
|
2482
|
+
launch_mul_mat_vec_q_moe<QK4_0, QI4_0, block_q4_0, VDR_Q4_0_Q8_1_MMVQ, vec_dot_q4_0_q8_1>(
|
|
2483
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2484
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2485
|
+
return true;
|
|
2486
|
+
case GGML_TYPE_Q4_1:
|
|
2487
|
+
launch_mul_mat_vec_q_moe<QK4_1, QI4_1, block_q4_1, VDR_Q4_1_Q8_1_MMVQ, vec_dot_q4_1_q8_1>(
|
|
2488
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2489
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2490
|
+
return true;
|
|
2491
|
+
case GGML_TYPE_Q5_0:
|
|
2492
|
+
launch_mul_mat_vec_q_moe<QK5_0, QI5_0, block_q5_0, VDR_Q5_0_Q8_1_MMVQ, vec_dot_q5_0_q8_1>(
|
|
2493
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2494
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2495
|
+
return true;
|
|
2496
|
+
case GGML_TYPE_Q5_1:
|
|
2497
|
+
launch_mul_mat_vec_q_moe<QK5_1, QI5_1, block_q5_1, VDR_Q5_1_Q8_1_MMVQ, vec_dot_q5_1_q8_1>(
|
|
2498
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2499
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2500
|
+
return true;
|
|
2501
|
+
case GGML_TYPE_Q8_0:
|
|
2502
|
+
launch_mul_mat_vec_q_moe<QK8_0, QI8_0, block_q8_0, VDR_Q8_0_Q8_1_MMVQ, vec_dot_q8_0_q8_1>(
|
|
2503
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2504
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2505
|
+
return true;
|
|
2506
|
+
case GGML_TYPE_Q2_K:
|
|
2507
|
+
launch_mul_mat_vec_q_moe<QK_K, QI2_K, block_q2_K, VDR_Q2_K_Q8_1_MMVQ, vec_dot_q2_K_q8_1>(
|
|
2508
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2509
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2510
|
+
return true;
|
|
2511
|
+
case GGML_TYPE_Q3_K:
|
|
2512
|
+
launch_mul_mat_vec_q_moe<QK_K, QI3_K, block_q3_K, VDR_Q3_K_Q8_1_MMVQ, vec_dot_q3_K_q8_1>(
|
|
2513
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2514
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2515
|
+
return true;
|
|
2516
|
+
case GGML_TYPE_Q4_K:
|
|
2517
|
+
launch_mul_mat_vec_q_moe<QK_K, QI4_K, block_q4_K, VDR_Q4_K_Q8_1_MMVQ, vec_dot_q4_K_q8_1>(
|
|
2518
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2519
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2520
|
+
return true;
|
|
2521
|
+
case GGML_TYPE_Q5_K:
|
|
2522
|
+
launch_mul_mat_vec_q_moe<QK_K, QI5_K, block_q5_K, VDR_Q5_K_Q8_1_MMVQ, vec_dot_q5_K_q8_1>(
|
|
2523
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2524
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2525
|
+
return true;
|
|
2526
|
+
case GGML_TYPE_Q6_K:
|
|
2527
|
+
launch_mul_mat_vec_q_moe<QK_K, QI6_K, block_q6_K, VDR_Q6_K_Q8_1_MMVQ, vec_dot_q6_K_q8_1>(
|
|
2528
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2529
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2530
|
+
return true;
|
|
2531
|
+
case GGML_TYPE_MXFP4:
|
|
2532
|
+
launch_mul_mat_vec_q_moe<QK_MXFP4, QI_MXFP4, block_mxfp4, VDR_MXFP4_Q8_1_MMVQ, vec_dot_mxfp4_q8_1>(
|
|
2533
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2534
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2535
|
+
return true;
|
|
2536
|
+
case GGML_TYPE_NVFP4:
|
|
2537
|
+
launch_mul_mat_vec_q_moe<QK_NVFP4, QI_NVFP4, block_nvfp4, VDR_NVFP4_Q8_1_MMVQ, vec_dot_nvfp4_q8_1>(
|
|
2538
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2539
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2540
|
+
return true;
|
|
2541
|
+
default:
|
|
2542
|
+
return false;
|
|
2543
|
+
}
|
|
2544
|
+
}
|
|
2545
|
+
|
|
2546
|
+
// Reorder (SoA) MoE expert GEMV: MoE expert/row/lane indexing (from mul_mat_vec_q_moe) with the
|
|
2547
|
+
// dense-reorder per-block reads (from mul_mat_vec_q_reorder). Each expert slice in vx_base is a
|
|
2548
|
+
// self-contained SoA, so nblocks = nrows*(ncols/qk) per expert and the constant expert stride holds.
|
|
2549
|
+
template <typename reorder_vec_dot_q_sycl>
|
|
2550
|
+
static void mul_mat_vec_q_moe_reorder(
|
|
2551
|
+
const void * __restrict__ vx_base, const void * __restrict__ vy_base,
|
|
2552
|
+
float * __restrict__ dst_base, const int32_t * __restrict__ ids_dev,
|
|
2553
|
+
const int ncols, const int nrows,
|
|
2554
|
+
const size_t expert_weight_stride, const size_t dst_row_stride,
|
|
2555
|
+
const size_t src1_row_stride,
|
|
2556
|
+
const sycl::nd_item<3> & item_ct1) {
|
|
2557
|
+
using block_type = ggml_sycl_reordered::block_q_t<reorder_vec_dot_q_sycl::gtype>;
|
|
2558
|
+
using block_traits = typename block_type::traits;
|
|
2559
|
+
|
|
2560
|
+
const int expert_idx = item_ct1.get_group(1);
|
|
2561
|
+
const int i02 = ids_dev[expert_idx];
|
|
2562
|
+
|
|
2563
|
+
const char * vx = (const char *) vx_base + (size_t) i02 * expert_weight_stride;
|
|
2564
|
+
const char * vy = (const char *) vy_base + (size_t) expert_idx * src1_row_stride;
|
|
2565
|
+
float * dst = (float *) ((char *) dst_base + (size_t) expert_idx * dst_row_stride);
|
|
2566
|
+
|
|
2567
|
+
const int row = item_ct1.get_group(2) * item_ct1.get_local_range(1) + item_ct1.get_local_id(1);
|
|
2568
|
+
if (row >= nrows) {
|
|
2569
|
+
return;
|
|
2570
|
+
}
|
|
2571
|
+
|
|
2572
|
+
const auto sg = item_ct1.get_sub_group();
|
|
2573
|
+
|
|
2574
|
+
const int blocks_per_row = ncols / block_traits::qk;
|
|
2575
|
+
constexpr int blocks_per_subgroup = ceil_div(block_traits::vdr_mmvq * WARP_SIZE, block_traits::qi);
|
|
2576
|
+
constexpr int block_elements_per_subgroup = block_traits::qi / block_traits::vdr_mmvq;
|
|
2577
|
+
const int nblocks = nrows * (ncols / block_traits::qk);
|
|
2578
|
+
|
|
2579
|
+
static_assert(blocks_per_subgroup > 0);
|
|
2580
|
+
static_assert(block_elements_per_subgroup > 0);
|
|
2581
|
+
|
|
2582
|
+
float partial_sum = 0.0f;
|
|
2583
|
+
for (int i = sg.get_local_linear_id() / block_elements_per_subgroup; i < blocks_per_row; i += blocks_per_subgroup) {
|
|
2584
|
+
const int ibx = row * blocks_per_row + i;
|
|
2585
|
+
|
|
2586
|
+
const auto bx_offset = block_type::get_block_offset(ibx, nblocks);
|
|
2587
|
+
const auto d_offset = block_type::get_d_offset(nrows, ncols, ibx);
|
|
2588
|
+
|
|
2589
|
+
const int iby = i * block_type::block_to_q8_1_ratio();
|
|
2590
|
+
const int8_t * q8_1_quant_ptr = (const int8_t *) vy + iby * QK8_1;
|
|
2591
|
+
const sycl::half2 * q8_1_ds_ptr = (const sycl::half2 *) ((const char *) vy + ncols + iby * sizeof(sycl::half2));
|
|
2592
|
+
|
|
2593
|
+
#pragma unroll
|
|
2594
|
+
for (int elem = 0; elem < block_elements_per_subgroup; elem += WARP_SIZE) {
|
|
2595
|
+
const int iqs = elem + block_traits::vdr_mmvq * (sg.get_local_linear_id() % block_elements_per_subgroup);
|
|
2596
|
+
partial_sum += reorder_vec_dot_q_sycl()(vx, bx_offset, d_offset, q8_1_quant_ptr, q8_1_ds_ptr, iqs);
|
|
2597
|
+
}
|
|
2598
|
+
}
|
|
2599
|
+
|
|
2600
|
+
auto sum = sycl::reduce_over_group(sg, partial_sum, std::plus<>());
|
|
2601
|
+
if (sg.leader()) {
|
|
2602
|
+
dst[row] = sum;
|
|
2603
|
+
}
|
|
2604
|
+
}
|
|
2605
|
+
|
|
2606
|
+
template <typename reorder_vec_dot_q_sycl>
|
|
2607
|
+
static void launch_mul_mat_vec_q_moe_reorder(
|
|
2608
|
+
const void * vx_base, const void * vy, const int32_t * ids_dev,
|
|
2609
|
+
float * dst_base, const int ncols, const int nrows, const int n_experts_used,
|
|
2610
|
+
const size_t expert_weight_stride, const size_t dst_row_stride,
|
|
2611
|
+
const size_t src1_row_stride,
|
|
2612
|
+
dpct::queue_ptr stream) {
|
|
2613
|
+
const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
|
|
2614
|
+
const sycl::range<3> block_nums(1, (unsigned) n_experts_used, (unsigned) block_num_y);
|
|
2615
|
+
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
|
|
2616
|
+
stream->submit([&](sycl::handler & cgh) {
|
|
2617
|
+
cgh.parallel_for(
|
|
2618
|
+
sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
|
2619
|
+
[=](sycl::nd_item<3> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
|
2620
|
+
mul_mat_vec_q_moe_reorder<reorder_vec_dot_q_sycl>(
|
|
2621
|
+
vx_base, vy, dst_base, ids_dev, ncols, nrows,
|
|
2622
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, item);
|
|
2623
|
+
});
|
|
2624
|
+
});
|
|
2625
|
+
}
|
|
2626
|
+
|
|
2627
|
+
bool ggml_sycl_mul_mat_vec_q_id_reorder(
|
|
2628
|
+
enum ggml_type src0_type,
|
|
2629
|
+
const void * vx_base,
|
|
2630
|
+
const void * vy,
|
|
2631
|
+
const int32_t * ids_dev,
|
|
2632
|
+
float * dst_base,
|
|
2633
|
+
int ncols,
|
|
2634
|
+
int nrows,
|
|
2635
|
+
int n_experts_used,
|
|
2636
|
+
size_t expert_weight_stride,
|
|
2637
|
+
size_t dst_row_stride,
|
|
2638
|
+
size_t src1_row_stride,
|
|
2639
|
+
dpct::queue_ptr stream) {
|
|
2640
|
+
switch (src0_type) {
|
|
2641
|
+
case GGML_TYPE_Q4_K:
|
|
2642
|
+
launch_mul_mat_vec_q_moe_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q4_K>>(
|
|
2643
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2644
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2645
|
+
return true;
|
|
2646
|
+
case GGML_TYPE_Q5_K:
|
|
2647
|
+
launch_mul_mat_vec_q_moe_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q5_K>>(
|
|
2648
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2649
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2650
|
+
return true;
|
|
2651
|
+
case GGML_TYPE_Q6_K:
|
|
2652
|
+
launch_mul_mat_vec_q_moe_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q6_K>>(
|
|
2653
|
+
vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
|
|
2654
|
+
expert_weight_stride, dst_row_stride, src1_row_stride, stream);
|
|
2655
|
+
return true;
|
|
2656
|
+
default:
|
|
2657
|
+
return false;
|
|
2658
|
+
}
|
|
2659
|
+
}
|