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.
Files changed (965) hide show
  1. checksums.yaml +4 -4
  2. data/.document +3 -0
  3. data/.rdoc_options +2 -0
  4. data/README.md +43 -9
  5. data/Rakefile +18 -3
  6. data/ext/dependencies.rb +10 -4
  7. data/ext/dependencies_for_windows.rb +17 -0
  8. data/ext/extconf.rb +20 -8
  9. data/ext/options.rb +54 -14
  10. data/ext/options_for_windows.rb +51 -0
  11. data/ext/ruby_whisper.c +35 -42
  12. data/ext/ruby_whisper.h +141 -0
  13. data/ext/ruby_whisper_context.c +157 -29
  14. data/ext/ruby_whisper_log_queue.c +180 -0
  15. data/ext/ruby_whisper_log_settable.h +46 -0
  16. data/ext/ruby_whisper_parakeet.c +49 -0
  17. data/ext/ruby_whisper_parakeet_context.c +304 -0
  18. data/ext/ruby_whisper_parakeet_context_params.c +117 -0
  19. data/ext/ruby_whisper_parakeet_model.c +84 -0
  20. data/ext/ruby_whisper_parakeet_params.c +548 -0
  21. data/ext/ruby_whisper_parakeet_segment.c +157 -0
  22. data/ext/ruby_whisper_parakeet_token.c +188 -0
  23. data/ext/ruby_whisper_parakeet_transcribe.cpp +58 -0
  24. data/ext/ruby_whisper_params.c +265 -73
  25. data/ext/ruby_whisper_segment.c +6 -6
  26. data/ext/ruby_whisper_transcribe.cpp +23 -15
  27. data/ext/ruby_whisper_vad_context.c +30 -10
  28. data/ext/ruby_whisper_vad_context_detect.cpp +8 -9
  29. data/ext/ruby_whisper_vad_params.c +4 -4
  30. data/ext/ruby_whisper_vad_segment.c +2 -2
  31. data/ext/sources/CMakeLists.txt +42 -3
  32. data/ext/sources/CMakePresets.json +95 -0
  33. data/ext/sources/cmake/parakeet-config.cmake.in +30 -0
  34. data/ext/sources/cmake/parakeet.pc.in +10 -0
  35. data/ext/sources/cmake/whisper.pc.in +2 -2
  36. data/ext/sources/examples/CMakeLists.txt +4 -2
  37. data/ext/sources/examples/bench/bench.cpp +1 -1
  38. data/ext/sources/examples/cli/cli.cpp +52 -10
  39. data/ext/sources/examples/common-ggml.cpp +4 -0
  40. data/ext/sources/examples/common-whisper.cpp +139 -67
  41. data/ext/sources/examples/common-whisper.h +11 -0
  42. data/ext/sources/examples/ffmpeg-transcode.cpp +211 -341
  43. data/ext/sources/examples/parakeet-cli/CMakeLists.txt +8 -0
  44. data/ext/sources/examples/parakeet-cli/parakeet-cli.cpp +243 -0
  45. data/ext/sources/examples/parakeet-quantize/CMakeLists.txt +7 -0
  46. data/ext/sources/examples/parakeet-quantize/parakeet-quantize.cpp +230 -0
  47. data/ext/sources/examples/server/server.cpp +199 -163
  48. data/ext/sources/examples/vad-speech-segments/speech.cpp +3 -2
  49. data/ext/sources/ggml/CMakeLists.txt +21 -14
  50. data/ext/sources/ggml/cmake/FindNCCL.cmake +36 -0
  51. data/ext/sources/ggml/cmake/ggml-config.cmake.in +12 -2
  52. data/ext/sources/ggml/include/ggml-alloc.h +1 -0
  53. data/ext/sources/ggml/include/ggml-backend.h +72 -10
  54. data/ext/sources/ggml/include/ggml-cuda.h +2 -2
  55. data/ext/sources/ggml/include/ggml-rpc.h +3 -3
  56. data/ext/sources/ggml/include/ggml-sycl.h +8 -0
  57. data/ext/sources/ggml/include/ggml.h +103 -9
  58. data/ext/sources/ggml/include/gguf.h +10 -2
  59. data/ext/sources/ggml/src/CMakeLists.txt +30 -6
  60. data/ext/sources/ggml/src/ggml-alloc.c +5 -1
  61. data/ext/sources/ggml/src/ggml-backend-impl.h +22 -2
  62. data/ext/sources/ggml/src/ggml-backend-meta.cpp +2266 -0
  63. data/ext/sources/ggml/src/ggml-backend-reg.cpp +12 -0
  64. data/ext/sources/ggml/src/ggml-backend.cpp +110 -9
  65. data/ext/sources/ggml/src/ggml-blas/ggml-blas.cpp +4 -0
  66. data/ext/sources/ggml/src/ggml-cann/aclnn_ops.cpp +672 -257
  67. data/ext/sources/ggml/src/ggml-cann/aclnn_ops.h +71 -0
  68. data/ext/sources/ggml/src/ggml-cann/common.h +20 -10
  69. data/ext/sources/ggml/src/ggml-cann/ggml-cann.cpp +211 -30
  70. data/ext/sources/ggml/src/ggml-common.h +24 -2
  71. data/ext/sources/ggml/src/ggml-cpu/CMakeLists.txt +59 -30
  72. data/ext/sources/ggml/src/ggml-cpu/amx/amx.cpp +2 -0
  73. data/ext/sources/ggml/src/ggml-cpu/amx/mmq.cpp +21 -22
  74. data/ext/sources/ggml/src/ggml-cpu/arch/arm/quants.c +194 -11
  75. data/ext/sources/ggml/src/ggml-cpu/arch/arm/repack.cpp +65 -0
  76. data/ext/sources/ggml/src/ggml-cpu/arch/loongarch/quants.c +151 -1
  77. data/ext/sources/ggml/src/ggml-cpu/arch/powerpc/quants.c +0 -1
  78. data/ext/sources/ggml/src/ggml-cpu/arch/riscv/quants.c +4279 -1292
  79. data/ext/sources/ggml/src/ggml-cpu/arch/riscv/repack.cpp +5 -35
  80. data/ext/sources/ggml/src/ggml-cpu/arch/s390/quants.c +0 -1
  81. data/ext/sources/ggml/src/ggml-cpu/arch/wasm/quants.c +72 -1
  82. data/ext/sources/ggml/src/ggml-cpu/arch/x86/quants.c +319 -31
  83. data/ext/sources/ggml/src/ggml-cpu/arch/x86/repack.cpp +1 -1
  84. data/ext/sources/ggml/src/ggml-cpu/arch-fallback.h +12 -2
  85. data/ext/sources/ggml/src/ggml-cpu/cmake/FindSMTIME.cmake +32 -0
  86. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu-impl.h +10 -0
  87. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.c +109 -5
  88. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.cpp +2 -0
  89. data/ext/sources/ggml/src/ggml-cpu/kleidiai/kleidiai.cpp +146 -134
  90. data/ext/sources/ggml/src/ggml-cpu/llamafile/sgemm.cpp +107 -82
  91. data/ext/sources/ggml/src/ggml-cpu/ops.cpp +501 -119
  92. data/ext/sources/ggml/src/ggml-cpu/ops.h +3 -0
  93. data/ext/sources/ggml/src/ggml-cpu/quants.c +106 -0
  94. data/ext/sources/ggml/src/ggml-cpu/quants.h +6 -0
  95. data/ext/sources/ggml/src/ggml-cpu/repack.cpp +3 -0
  96. data/ext/sources/ggml/src/ggml-cpu/simd-gemm.h +91 -1
  97. data/ext/sources/ggml/src/ggml-cpu/simd-mappings.h +14 -16
  98. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime.cpp +1402 -687
  99. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime.h +8 -0
  100. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime1_kernels.cpp +597 -2766
  101. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime2_kernels.cpp +5768 -0
  102. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_env.cpp +320 -0
  103. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_env.h +55 -0
  104. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_kernels.h +182 -19
  105. data/ext/sources/ggml/src/ggml-cpu/spacemit/repack.cpp +1795 -0
  106. data/ext/sources/ggml/src/ggml-cpu/spacemit/repack.h +14 -0
  107. data/ext/sources/ggml/src/ggml-cpu/spacemit/rvv_kernels.cpp +3178 -0
  108. data/ext/sources/ggml/src/ggml-cpu/spacemit/rvv_kernels.h +95 -0
  109. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_barrier.h +34 -0
  110. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_mem_pool.cpp +760 -0
  111. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_mem_pool.h +32 -0
  112. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_tcm.h +409 -0
  113. data/ext/sources/ggml/src/ggml-cpu/vec.cpp +39 -55
  114. data/ext/sources/ggml/src/ggml-cpu/vec.h +225 -240
  115. data/ext/sources/ggml/src/ggml-cuda/CMakeLists.txt +17 -7
  116. data/ext/sources/ggml/src/ggml-cuda/allreduce.cu +971 -0
  117. data/ext/sources/ggml/src/ggml-cuda/allreduce.cuh +29 -0
  118. data/ext/sources/ggml/src/ggml-cuda/argsort.cu +62 -26
  119. data/ext/sources/ggml/src/ggml-cuda/binbcast.cu +134 -64
  120. data/ext/sources/ggml/src/ggml-cuda/binbcast.cuh +1 -0
  121. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cu +81 -0
  122. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cuh +3 -0
  123. data/ext/sources/ggml/src/ggml-cuda/common.cuh +246 -28
  124. data/ext/sources/ggml/src/ggml-cuda/concat.cu +134 -116
  125. data/ext/sources/ggml/src/ggml-cuda/conv-transpose-1d.cu +14 -12
  126. data/ext/sources/ggml/src/ggml-cuda/conv2d-transpose.cu +45 -21
  127. data/ext/sources/ggml/src/ggml-cuda/conv2d-transpose.cuh +1 -0
  128. data/ext/sources/ggml/src/ggml-cuda/convert.cu +139 -34
  129. data/ext/sources/ggml/src/ggml-cuda/convert.cuh +10 -0
  130. data/ext/sources/ggml/src/ggml-cuda/cpy.cu +88 -29
  131. data/ext/sources/ggml/src/ggml-cuda/dequantize.cuh +22 -0
  132. data/ext/sources/ggml/src/ggml-cuda/fattn-common.cuh +287 -49
  133. data/ext/sources/ggml/src/ggml-cuda/fattn-mma-f16.cuh +335 -130
  134. data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cu +12 -0
  135. data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cuh +127 -24
  136. data/ext/sources/ggml/src/ggml-cuda/fattn-vec.cuh +40 -15
  137. data/ext/sources/ggml/src/ggml-cuda/fattn-wmma-f16.cu +18 -9
  138. data/ext/sources/ggml/src/ggml-cuda/fattn.cu +169 -60
  139. data/ext/sources/ggml/src/ggml-cuda/fattn.cuh +2 -0
  140. data/ext/sources/ggml/src/ggml-cuda/fwht.cu +101 -0
  141. data/ext/sources/ggml/src/ggml-cuda/fwht.cuh +4 -0
  142. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cu +109 -45
  143. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cuh +10 -0
  144. data/ext/sources/ggml/src/ggml-cuda/getrows.cu +48 -23
  145. data/ext/sources/ggml/src/ggml-cuda/ggml-cuda.cu +2034 -2104
  146. data/ext/sources/ggml/src/ggml-cuda/im2col.cu +32 -29
  147. data/ext/sources/ggml/src/ggml-cuda/mean.cu +4 -2
  148. data/ext/sources/ggml/src/ggml-cuda/mma.cuh +242 -195
  149. data/ext/sources/ggml/src/ggml-cuda/mmf.cuh +3 -3
  150. data/ext/sources/ggml/src/ggml-cuda/mmq.cu +25 -12
  151. data/ext/sources/ggml/src/ggml-cuda/mmq.cuh +502 -423
  152. data/ext/sources/ggml/src/ggml-cuda/mmvf.cu +19 -12
  153. data/ext/sources/ggml/src/ggml-cuda/mmvq.cu +562 -97
  154. data/ext/sources/ggml/src/ggml-cuda/mmvq.cuh +6 -1
  155. data/ext/sources/ggml/src/ggml-cuda/norm.cu +36 -10
  156. data/ext/sources/ggml/src/ggml-cuda/out-prod.cu +66 -7
  157. data/ext/sources/ggml/src/ggml-cuda/quantize.cu +133 -26
  158. data/ext/sources/ggml/src/ggml-cuda/quantize.cuh +1 -1
  159. data/ext/sources/ggml/src/ggml-cuda/reduce_rows.cuh +5 -1
  160. data/ext/sources/ggml/src/ggml-cuda/rope.cu +11 -4
  161. data/ext/sources/ggml/src/ggml-cuda/scale.cu +4 -1
  162. data/ext/sources/ggml/src/ggml-cuda/set-rows.cu +78 -10
  163. data/ext/sources/ggml/src/ggml-cuda/snake.cu +72 -0
  164. data/ext/sources/ggml/src/ggml-cuda/snake.cuh +8 -0
  165. data/ext/sources/ggml/src/ggml-cuda/softcap.cu +4 -1
  166. data/ext/sources/ggml/src/ggml-cuda/ssm-conv.cu +45 -13
  167. data/ext/sources/ggml/src/ggml-cuda/ssm-conv.cuh +1 -1
  168. data/ext/sources/ggml/src/ggml-cuda/ssm-scan.cu +40 -18
  169. data/ext/sources/ggml/src/ggml-cuda/sumrows.cu +8 -4
  170. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_16.cu +1 -0
  171. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_32.cu +1 -0
  172. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_8.cu +2 -0
  173. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_2.cu +1 -0
  174. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_4.cu +1 -0
  175. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_16.cu +1 -0
  176. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_32.cu +1 -0
  177. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_4.cu +1 -0
  178. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_8.cu +2 -0
  179. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_32-ncols2_2.cu +1 -0
  180. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_16.cu +1 -0
  181. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_2.cu +1 -0
  182. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_4.cu +1 -0
  183. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_8.cu +2 -0
  184. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_2.cu +1 -0
  185. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_4.cu +1 -0
  186. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_8.cu +2 -0
  187. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq192-dv128.cu +5 -0
  188. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq320-dv256.cu +5 -0
  189. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq512-dv512.cu +5 -0
  190. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-bf16.cu +7 -0
  191. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-f16.cu +7 -0
  192. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q4_0.cu +7 -0
  193. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q4_1.cu +7 -0
  194. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q5_0.cu +7 -0
  195. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q5_1.cu +7 -0
  196. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q8_0.cu +7 -0
  197. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-f16-bf16.cu +7 -0
  198. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q4_0-bf16.cu +7 -0
  199. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q4_1-bf16.cu +7 -0
  200. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q5_0-bf16.cu +7 -0
  201. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q5_1-bf16.cu +7 -0
  202. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q8_0-bf16.cu +7 -0
  203. data/ext/sources/ggml/src/ggml-cuda/template-instances/mmq-instance-nvfp4.cu +5 -0
  204. data/ext/sources/ggml/src/ggml-cuda/template-instances/mmq-instance-q1_0.cu +5 -0
  205. data/ext/sources/ggml/src/ggml-cuda/top-k.cu +5 -4
  206. data/ext/sources/ggml/src/ggml-cuda/topk-moe.cu +33 -24
  207. data/ext/sources/ggml/src/ggml-cuda/unary.cu +31 -2
  208. data/ext/sources/ggml/src/ggml-cuda/unary.cuh +2 -0
  209. data/ext/sources/ggml/src/ggml-cuda/vecdotq.cuh +80 -0
  210. data/ext/sources/ggml/src/ggml-cuda/vendors/cuda.h +7 -2
  211. data/ext/sources/ggml/src/ggml-cuda/vendors/hip.h +23 -4
  212. data/ext/sources/ggml/src/ggml-cuda/vendors/musa.h +4 -0
  213. data/ext/sources/ggml/src/ggml-hexagon/CMakeLists.txt +1 -5
  214. data/ext/sources/ggml/src/ggml-hexagon/ggml-hexagon.cpp +2788 -1762
  215. data/ext/sources/ggml/src/ggml-hexagon/htp/CMakeLists.txt +13 -4
  216. data/ext/sources/ggml/src/ggml-hexagon/htp/act-ops.c +53 -84
  217. data/ext/sources/ggml/src/ggml-hexagon/htp/argsort-ops.c +25 -12
  218. data/ext/sources/ggml/src/ggml-hexagon/htp/binary-ops.c +165 -184
  219. data/ext/sources/ggml/src/ggml-hexagon/htp/cmake-toolchain.cmake +17 -19
  220. data/ext/sources/ggml/src/ggml-hexagon/htp/concat-ops.c +277 -0
  221. data/ext/sources/ggml/src/ggml-hexagon/htp/cpy-ops.c +170 -127
  222. data/ext/sources/ggml/src/ggml-hexagon/htp/cumsum-ops.c +270 -0
  223. data/ext/sources/ggml/src/ggml-hexagon/htp/diag-ops.c +216 -0
  224. data/ext/sources/ggml/src/ggml-hexagon/htp/fill-ops.c +123 -0
  225. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +1774 -396
  226. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.h +303 -0
  227. data/ext/sources/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c +1148 -0
  228. data/ext/sources/ggml/src/ggml-hexagon/htp/get-rows-ops.c +148 -42
  229. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-common.h +80 -0
  230. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.c +2 -2
  231. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.h +255 -62
  232. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dump.h +9 -0
  233. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-profile.h +64 -0
  234. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-utils.h +25 -21
  235. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h +555 -0
  236. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h +1303 -0
  237. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.c +167 -0
  238. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.h +157 -0
  239. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-utils.h +222 -0
  240. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ctx.h +104 -13
  241. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ops.h +222 -57
  242. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-vtcm.h +19 -0
  243. data/ext/sources/ggml/src/ggml-hexagon/htp/htp_iface.idl +10 -3
  244. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-base.h +78 -26
  245. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-copy.h +27 -10
  246. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-div.h +63 -23
  247. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-exp.h +48 -8
  248. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h +232 -0
  249. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-flash-attn.h +47 -0
  250. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-log.h +65 -0
  251. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-flat.h +1511 -0
  252. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h +1200 -0
  253. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-pow.h +42 -0
  254. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-repl.h +74 -0
  255. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h +40 -0
  256. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sin-cos.h +90 -0
  257. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-utils.h +5 -8
  258. data/ext/sources/ggml/src/ggml-hexagon/htp/main.c +625 -816
  259. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.c +3052 -2166
  260. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.h +650 -0
  261. data/ext/sources/ggml/src/ggml-hexagon/htp/pad-ops.c +547 -0
  262. data/ext/sources/ggml/src/ggml-hexagon/htp/repeat-ops.c +148 -0
  263. data/ext/sources/ggml/src/ggml-hexagon/htp/rope-ops.c +337 -106
  264. data/ext/sources/ggml/src/ggml-hexagon/htp/set-rows-ops.c +59 -37
  265. data/ext/sources/ggml/src/ggml-hexagon/htp/softmax-ops.c +121 -133
  266. data/ext/sources/ggml/src/ggml-hexagon/htp/solve-tri-ops.c +267 -0
  267. data/ext/sources/ggml/src/ggml-hexagon/htp/ssm-conv.c +245 -151
  268. data/ext/sources/ggml/src/ggml-hexagon/htp/sum-rows-ops.c +6 -6
  269. data/ext/sources/ggml/src/ggml-hexagon/htp/unary-ops.c +719 -45
  270. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.c +15 -3
  271. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.h +8 -0
  272. data/ext/sources/ggml/src/ggml-hexagon/htp-opnode.h +390 -0
  273. data/ext/sources/ggml/src/ggml-hexagon/libggml-htp.inf +3 -5
  274. data/ext/sources/ggml/src/ggml-hip/CMakeLists.txt +27 -9
  275. data/ext/sources/ggml/src/ggml-impl.h +6 -1
  276. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.cpp +207 -18
  277. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.h +36 -2
  278. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.m +186 -29
  279. data/ext/sources/ggml/src/ggml-metal/ggml-metal-impl.h +118 -0
  280. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.cpp +322 -21
  281. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.h +4 -0
  282. data/ext/sources/ggml/src/ggml-metal/ggml-metal.cpp +39 -26
  283. data/ext/sources/ggml/src/ggml-metal/ggml-metal.metal +1226 -467
  284. data/ext/sources/ggml/src/ggml-musa/CMakeLists.txt +5 -6
  285. data/ext/sources/ggml/src/ggml-opencl/CMakeLists.txt +67 -5
  286. data/ext/sources/ggml/src/ggml-opencl/fa_tune.h +92 -0
  287. data/ext/sources/ggml/src/ggml-opencl/ggml-opencl.cpp +16290 -6246
  288. data/ext/sources/ggml/src/ggml-opencl/kernels/concat.cl +67 -0
  289. data/ext/sources/ggml/src/ggml-opencl/kernels/cpy.cl +59 -0
  290. data/ext/sources/ggml/src/ggml-opencl/kernels/cvt.cl +1997 -92
  291. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f16.cl +81 -41
  292. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32.cl +88 -39
  293. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl +1995 -96
  294. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q4_0.cl +1615 -0
  295. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q8_0.cl +1486 -0
  296. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_pre_f16.cl +156 -0
  297. data/ext/sources/ggml/src/ggml-opencl/kernels/gated_delta_net.cl +249 -0
  298. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_mxfp4_f32_ns.cl +374 -0
  299. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_0_f32_ns.cl +324 -0
  300. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_1_f32_ns.cl +326 -0
  301. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_k_f32_ns.cl +348 -0
  302. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_0_f32_ns.cl +328 -0
  303. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_1_f32_ns.cl +330 -0
  304. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_k_f32_ns.cl +356 -0
  305. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q6_k_f32_ns.cl +335 -0
  306. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_iq4_nl_f32.cl +150 -0
  307. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q1_0_f32.cl +94 -0
  308. data/ext/sources/ggml/src/ggml-opencl/kernels/{mul_mat_Ab_Bi_8x4.cl → gemm_noshuffle_q4_0_f32.cl} +1 -1
  309. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q4_k_f32.cl +172 -0
  310. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_0_f32.cl +131 -0
  311. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_1_f32.cl +134 -0
  312. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_k_f32.cl +176 -0
  313. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q6_k_f32.cl +140 -0
  314. data/ext/sources/ggml/src/ggml-opencl/kernels/{mul_mm_q8_0_f32_8x4.cl → gemm_noshuffle_q8_0_f32.cl} +1 -1
  315. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_xmem_f16_f32_os8.cl +233 -0
  316. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_mxfp4_f32_ns.cl +165 -0
  317. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_0_f32_ns.cl +120 -0
  318. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_1_f32_ns.cl +123 -0
  319. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_k_f32_ns.cl +155 -0
  320. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_0_f32_ns.cl +123 -0
  321. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_1_f32_ns.cl +125 -0
  322. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_k_f32_ns.cl +160 -0
  323. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q6_k_f32_ns.cl +141 -0
  324. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_iq4_nl_f32.cl +302 -0
  325. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q1_0_f32.cl +121 -0
  326. data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle_general.cl → gemv_noshuffle_q4_0_f32.cl} +5 -5
  327. data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle.cl → gemv_noshuffle_q4_0_f32_spec.cl} +5 -5
  328. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32.cl +318 -0
  329. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_0_f32.cl +291 -0
  330. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_1_f32.cl +294 -0
  331. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_k_f32.cl +326 -0
  332. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32.cl +293 -0
  333. data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle_general_q8_0_f32.cl → gemv_noshuffle_q8_0_f32.cl} +1 -1
  334. data/ext/sources/ggml/src/ggml-opencl/kernels/get_rows.cl +15 -9
  335. data/ext/sources/ggml/src/ggml-opencl/kernels/moe_reorder_b.cl +30 -0
  336. data/ext/sources/ggml/src/ggml-opencl/kernels/moe_sort_by_expert.cl +82 -0
  337. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_iq4_nl_f32_l4_lm.cl +171 -0
  338. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q1_0_f32_l4_lm.cl +156 -0
  339. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q4_k_f32_l4_lm.cl +179 -0
  340. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_0_f32_l4_lm.cl +173 -0
  341. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_1_f32_l4_lm.cl +175 -0
  342. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_k_f32_l4_lm.cl +192 -0
  343. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_l4.cl +1149 -0
  344. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_iq4_nl_f32.cl +164 -0
  345. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_iq4_nl_f32_flat.cl +202 -0
  346. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32.cl +141 -0
  347. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32_flat.cl +190 -0
  348. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q4_k_f32_flat.cl +196 -0
  349. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_0_f32.cl +241 -0
  350. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_0_f32_flat.cl +243 -0
  351. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_1_f32.cl +243 -0
  352. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_1_f32_flat.cl +247 -0
  353. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32.cl +187 -0
  354. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32_flat.cl +203 -0
  355. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q6_k_f32_flat.cl +48 -64
  356. data/ext/sources/ggml/src/ggml-opencl/kernels/norm.cl +5 -2
  357. data/ext/sources/ggml/src/ggml-opencl/kernels/set_rows.cl +500 -0
  358. data/ext/sources/ggml/src/ggml-opencl/libdl.h +79 -0
  359. data/ext/sources/ggml/src/ggml-openvino/.clang-format +0 -5
  360. data/ext/sources/ggml/src/ggml-openvino/CMakeLists.txt +2 -4
  361. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.cpp +740 -127
  362. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.h +76 -23
  363. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.cpp +75 -14
  364. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.h +29 -8
  365. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino.cpp +339 -69
  366. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.cpp +330 -192
  367. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.h +10 -4
  368. data/ext/sources/ggml/src/ggml-openvino/openvino/decoder.h +56 -16
  369. data/ext/sources/ggml/src/ggml-openvino/openvino/frontend.h +1 -1
  370. data/ext/sources/ggml/src/ggml-openvino/openvino/input_model.h +4 -4
  371. data/ext/sources/ggml/src/ggml-openvino/openvino/node_context.h +94 -37
  372. data/ext/sources/ggml/src/ggml-openvino/openvino/op/add_id.cpp +76 -0
  373. data/ext/sources/ggml/src/ggml-openvino/openvino/op/argsort.cpp +47 -0
  374. data/ext/sources/ggml/src/ggml-openvino/openvino/op/clamp.cpp +33 -0
  375. data/ext/sources/ggml/src/ggml-openvino/openvino/op/concat.cpp +48 -0
  376. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cont.cpp +8 -16
  377. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cpy.cpp +14 -1
  378. data/ext/sources/ggml/src/ggml-openvino/openvino/op/div.cpp +146 -0
  379. data/ext/sources/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp +108 -21
  380. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp +282 -0
  381. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.hpp +65 -0
  382. data/ext/sources/ggml/src/ggml-openvino/openvino/op/get_rows.cpp +2 -9
  383. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_geglu.cpp +21 -7
  384. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp +41 -8
  385. data/ext/sources/ggml/src/ggml-openvino/openvino/op/im2col.cpp +120 -0
  386. data/ext/sources/ggml/src/ggml-openvino/openvino/op/l2_norm.cpp +44 -0
  387. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mul_mat_id.cpp +226 -0
  388. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mulmat.cpp +19 -9
  389. data/ext/sources/ggml/src/ggml-openvino/openvino/op/norm.cpp +58 -0
  390. data/ext/sources/ggml/src/ggml-openvino/openvino/op/pad.cpp +95 -0
  391. data/ext/sources/ggml/src/ggml-openvino/openvino/op/permute.cpp +58 -13
  392. data/ext/sources/ggml/src/ggml-openvino/openvino/op/repeat.cpp +74 -0
  393. data/ext/sources/ggml/src/ggml-openvino/openvino/op/reshape.cpp +13 -6
  394. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rms_norm.cpp +1 -1
  395. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rope.cpp +161 -39
  396. data/ext/sources/ggml/src/ggml-openvino/openvino/op/set_rows.cpp +3 -3
  397. data/ext/sources/ggml/src/ggml-openvino/openvino/op/softmax.cpp +126 -49
  398. data/ext/sources/ggml/src/ggml-openvino/openvino/op/ssm_conv.cpp +59 -0
  399. data/ext/sources/ggml/src/ggml-openvino/openvino/op/sum_rows.cpp +27 -0
  400. data/ext/sources/ggml/src/ggml-openvino/openvino/op/transpose.cpp +32 -1
  401. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_silu.cpp +1 -1
  402. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_softplus.cpp +38 -0
  403. data/ext/sources/ggml/src/ggml-openvino/openvino/op/view.cpp +90 -25
  404. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.cpp +41 -22
  405. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.h +18 -4
  406. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/mark_decompression_convert_constant_folding.h +1 -1
  407. data/ext/sources/ggml/src/ggml-openvino/openvino/rt_info/weightless_caching_attributes.hpp +41 -0
  408. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.cpp +70 -43
  409. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.h +5 -4
  410. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.cpp +612 -36
  411. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.h +29 -26
  412. data/ext/sources/ggml/src/ggml-openvino/utils.cpp +460 -114
  413. data/ext/sources/ggml/src/ggml-openvino/utils.h +32 -9
  414. data/ext/sources/ggml/src/ggml-opt.cpp +1 -0
  415. data/ext/sources/ggml/src/ggml-quants.c +365 -114
  416. data/ext/sources/ggml/src/ggml-quants.h +6 -0
  417. data/ext/sources/ggml/src/ggml-rpc/CMakeLists.txt +24 -0
  418. data/ext/sources/ggml/src/ggml-rpc/ggml-rpc.cpp +167 -311
  419. data/ext/sources/ggml/src/ggml-rpc/transport.cpp +683 -0
  420. data/ext/sources/ggml/src/ggml-rpc/transport.h +34 -0
  421. data/ext/sources/ggml/src/ggml-sycl/CMakeLists.txt +50 -4
  422. data/ext/sources/ggml/src/ggml-sycl/add-id.cpp +1 -1
  423. data/ext/sources/ggml/src/ggml-sycl/backend.hpp +5 -1
  424. data/ext/sources/ggml/src/ggml-sycl/binbcast.cpp +12 -0
  425. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.cpp +102 -0
  426. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.hpp +8 -0
  427. data/ext/sources/ggml/src/ggml-sycl/common.cpp +72 -2
  428. data/ext/sources/ggml/src/ggml-sycl/common.hpp +59 -2
  429. data/ext/sources/ggml/src/ggml-sycl/concat.cpp +21 -1
  430. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.cpp +158 -0
  431. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.hpp +10 -0
  432. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.cpp +125 -0
  433. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.hpp +10 -0
  434. data/ext/sources/ggml/src/ggml-sycl/conv2d.cpp +150 -0
  435. data/ext/sources/ggml/src/ggml-sycl/conv2d.hpp +10 -0
  436. data/ext/sources/ggml/src/ggml-sycl/conv3d.cpp +224 -0
  437. data/ext/sources/ggml/src/ggml-sycl/conv3d.hpp +8 -0
  438. data/ext/sources/ggml/src/ggml-sycl/convert.cpp +121 -13
  439. data/ext/sources/ggml/src/ggml-sycl/convert.hpp +9 -0
  440. data/ext/sources/ggml/src/ggml-sycl/cpy.cpp +706 -0
  441. data/ext/sources/ggml/src/ggml-sycl/cpy.hpp +281 -0
  442. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.cpp +255 -0
  443. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.hpp +7 -0
  444. data/ext/sources/ggml/src/ggml-sycl/cumsum.cpp +148 -0
  445. data/ext/sources/ggml/src/ggml-sycl/cumsum.hpp +5 -0
  446. data/ext/sources/ggml/src/ggml-sycl/dequantize.hpp +678 -0
  447. data/ext/sources/ggml/src/ggml-sycl/diag.cpp +67 -0
  448. data/ext/sources/ggml/src/ggml-sycl/diag.hpp +5 -0
  449. data/ext/sources/ggml/src/ggml-sycl/dmmv.cpp +997 -244
  450. data/ext/sources/ggml/src/ggml-sycl/dpct/helper.hpp +15 -7
  451. data/ext/sources/ggml/src/ggml-sycl/element_wise.cpp +215 -204
  452. data/ext/sources/ggml/src/ggml-sycl/element_wise.hpp +2 -2
  453. data/ext/sources/ggml/src/ggml-sycl/fattn-buffers.cpp +56 -0
  454. data/ext/sources/ggml/src/ggml-sycl/fattn-buffers.hpp +63 -0
  455. data/ext/sources/ggml/src/ggml-sycl/fattn-common.hpp +7 -5
  456. data/ext/sources/ggml/src/ggml-sycl/fattn-tile.cpp +4 -0
  457. data/ext/sources/ggml/src/ggml-sycl/fattn-tile.hpp +76 -168
  458. data/ext/sources/ggml/src/ggml-sycl/fattn-vec.hpp +7 -0
  459. data/ext/sources/ggml/src/ggml-sycl/fattn.cpp +3 -1
  460. data/ext/sources/ggml/src/ggml-sycl/fill.cpp +55 -0
  461. data/ext/sources/ggml/src/ggml-sycl/fill.hpp +5 -0
  462. data/ext/sources/ggml/src/ggml-sycl/gated_delta_net.cpp +69 -31
  463. data/ext/sources/ggml/src/ggml-sycl/gated_delta_net.hpp +1 -0
  464. data/ext/sources/ggml/src/ggml-sycl/gemm.hpp +3 -0
  465. data/ext/sources/ggml/src/ggml-sycl/getrows.cpp +79 -3
  466. data/ext/sources/ggml/src/ggml-sycl/ggml-sycl.cpp +1758 -455
  467. data/ext/sources/ggml/src/ggml-sycl/im2col.cpp +353 -89
  468. data/ext/sources/ggml/src/ggml-sycl/im2col.hpp +5 -3
  469. data/ext/sources/ggml/src/ggml-sycl/mmvq.cpp +1542 -39
  470. data/ext/sources/ggml/src/ggml-sycl/mmvq.hpp +33 -0
  471. data/ext/sources/ggml/src/ggml-sycl/norm.cpp +103 -49
  472. data/ext/sources/ggml/src/ggml-sycl/outprod.cpp +45 -9
  473. data/ext/sources/ggml/src/ggml-sycl/pad.cpp +27 -27
  474. data/ext/sources/ggml/src/ggml-sycl/pool.cpp +185 -0
  475. data/ext/sources/ggml/src/ggml-sycl/pool.hpp +22 -0
  476. data/ext/sources/ggml/src/ggml-sycl/presets.hpp +3 -1
  477. data/ext/sources/ggml/src/ggml-sycl/quants.hpp +71 -0
  478. data/ext/sources/ggml/src/ggml-sycl/set_rows.cpp +17 -3
  479. data/ext/sources/ggml/src/ggml-sycl/softmax.cpp +9 -10
  480. data/ext/sources/ggml/src/ggml-sycl/solve_tri.cpp +172 -0
  481. data/ext/sources/ggml/src/ggml-sycl/solve_tri.hpp +8 -0
  482. data/ext/sources/ggml/src/ggml-sycl/ssm_conv.cpp +6 -1
  483. data/ext/sources/ggml/src/ggml-sycl/ssm_scan.cpp +156 -0
  484. data/ext/sources/ggml/src/ggml-sycl/ssm_scan.hpp +5 -0
  485. data/ext/sources/ggml/src/ggml-sycl/sycl_hw.cpp +62 -10
  486. data/ext/sources/ggml/src/ggml-sycl/sycl_hw.hpp +18 -6
  487. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-tile-instance-dkq512-dv512.cpp +6 -0
  488. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-f16.cpp +1 -0
  489. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q4_0.cpp +1 -0
  490. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q4_1.cpp +1 -0
  491. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q5_0.cpp +1 -0
  492. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q5_1.cpp +1 -0
  493. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q8_0.cpp +1 -0
  494. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-f16.cpp +1 -0
  495. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q4_0.cpp +1 -0
  496. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q4_1.cpp +1 -0
  497. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q5_0.cpp +1 -0
  498. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q5_1.cpp +1 -0
  499. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q8_0.cpp +1 -0
  500. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-f16.cpp +1 -0
  501. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q4_0.cpp +1 -0
  502. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q4_1.cpp +1 -0
  503. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q5_0.cpp +1 -0
  504. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q5_1.cpp +1 -0
  505. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q8_0.cpp +1 -0
  506. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-f16.cpp +1 -0
  507. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q4_0.cpp +1 -0
  508. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q4_1.cpp +1 -0
  509. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q5_0.cpp +1 -0
  510. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q5_1.cpp +1 -0
  511. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q8_0.cpp +1 -0
  512. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-f16.cpp +1 -0
  513. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q4_0.cpp +1 -0
  514. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q4_1.cpp +1 -0
  515. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q5_0.cpp +1 -0
  516. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q5_1.cpp +1 -0
  517. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q8_0.cpp +1 -0
  518. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-f16.cpp +1 -0
  519. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q4_0.cpp +1 -0
  520. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q4_1.cpp +1 -0
  521. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q5_0.cpp +1 -0
  522. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q5_1.cpp +1 -0
  523. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q8_0.cpp +1 -0
  524. data/ext/sources/ggml/src/ggml-sycl/type.hpp +112 -0
  525. data/ext/sources/ggml/src/ggml-sycl/upscale.cpp +410 -0
  526. data/ext/sources/ggml/src/ggml-sycl/upscale.hpp +9 -0
  527. data/ext/sources/ggml/src/ggml-sycl/vecdotq.hpp +242 -45
  528. data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend-buffer.cpp +4 -0
  529. data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend-device.cpp +2 -0
  530. data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend.cpp +2 -0
  531. data/ext/sources/ggml/src/ggml-virtgpu/virtgpu-shm.cpp +1 -0
  532. data/ext/sources/ggml/src/ggml-virtgpu/virtgpu.cpp +1 -0
  533. data/ext/sources/ggml/src/ggml-virtgpu/virtgpu.h +0 -2
  534. data/ext/sources/ggml/src/ggml-vulkan/CMakeLists.txt +16 -0
  535. data/ext/sources/ggml/src/ggml-vulkan/ggml-vulkan.cpp +2843 -700
  536. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/CMakeLists.txt +4 -0
  537. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/col2im_1d.comp +61 -0
  538. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/contig_copy.comp +6 -2
  539. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp +146 -13
  540. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp +431 -0
  541. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy.comp +3 -1
  542. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy_from_quant.comp +1 -1
  543. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy_to_quant.comp +25 -1
  544. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl +88 -0
  545. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl +643 -1
  546. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_nvfp4.comp +32 -0
  547. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q1_0.comp +29 -0
  548. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/diag.comp +3 -4
  549. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dot_product_funcs.glsl +27 -0
  550. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/feature-tests/coopmat2_decode_vector.comp +7 -0
  551. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +198 -48
  552. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl +60 -59
  553. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +116 -113
  554. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp +122 -31
  555. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_dequant.glsl +131 -0
  556. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_mmq_funcs.glsl +203 -0
  557. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/fwht.comp +115 -0
  558. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gated_delta_net.comp +125 -64
  559. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_binary_head.glsl +0 -1
  560. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl +21 -19
  561. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/get_rows_back.comp +25 -0
  562. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl +29 -1
  563. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_main.glsl +17 -11
  564. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/im2col.comp +76 -54
  565. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/im2col_3d.comp +0 -1
  566. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/l2_norm.comp +4 -7
  567. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/log.comp +0 -1
  568. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec.comp +122 -27
  569. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iface.glsl +6 -6
  570. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q2_k.comp +1 -1
  571. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q4_k.comp +1 -1
  572. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q5_k.comp +1 -1
  573. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp +22 -24
  574. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_funcs.glsl +88 -55
  575. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +42 -40
  576. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +49 -15
  577. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +222 -171
  578. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_funcs.glsl +8 -8
  579. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_shmem_types.glsl +24 -9
  580. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/multi_add.comp +0 -1
  581. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/norm.comp +10 -10
  582. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/repeat_back.comp +3 -3
  583. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/roll.comp +3 -3
  584. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_funcs.glsl +5 -2
  585. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_head.glsl +0 -1
  586. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_params.glsl +3 -2
  587. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/snake.comp +49 -0
  588. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ssm_conv.comp +11 -1
  589. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tri.comp +3 -4
  590. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl +79 -2
  591. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/unary.comp +168 -0
  592. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +282 -211
  593. data/ext/sources/ggml/src/ggml-webgpu/CMakeLists.txt +5 -2
  594. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +2209 -283
  595. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu.cpp +2618 -1416
  596. data/ext/sources/ggml/src/ggml-webgpu/pre_wgsl.hpp +37 -7
  597. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/add_id.wgsl +64 -0
  598. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/binary.wgsl +8 -7
  599. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl +90 -95
  600. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl +19 -1
  601. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/conv2d.wgsl +165 -0
  602. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{cpy.tmpl.wgsl → cpy.wgsl} +25 -50
  603. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl +107 -184
  604. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_quant_staging.tmpl +124 -0
  605. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_tile.wgsl +397 -0
  606. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_blk.wgsl +101 -0
  607. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_reduce.wgsl +84 -0
  608. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl +619 -0
  609. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/gated_delta_net.wgsl +149 -0
  610. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl +204 -78
  611. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/glu.wgsl +155 -0
  612. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/im2col.wgsl +101 -0
  613. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl +805 -526
  614. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id.wgsl +195 -0
  615. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_gather.wgsl +52 -0
  616. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_vec.wgsl +154 -0
  617. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_reg_tile.wgsl +8 -6
  618. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_subgroup_matrix.wgsl +5 -1
  619. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl +90 -413
  620. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl +1553 -0
  621. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl +297 -0
  622. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quant_inner_loops.tmpl +21 -0
  623. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl +178 -0
  624. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm_mul.wgsl +152 -0
  625. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{rope.tmpl.wgsl → rope.wgsl} +71 -142
  626. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/row_norm.wgsl +153 -0
  627. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/scale.wgsl +6 -4
  628. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set.wgsl +109 -0
  629. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set_rows.wgsl +2 -3
  630. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set_rows_quant.wgsl +224 -0
  631. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{soft_max.tmpl.wgsl → soft_max.wgsl} +106 -206
  632. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/solve_tri.wgsl +121 -0
  633. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/ssm_conv.wgsl +65 -0
  634. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl +193 -0
  635. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/unary.wgsl +68 -48
  636. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/upscale.wgsl +240 -0
  637. data/ext/sources/ggml/src/ggml-zdnn/ggml-zdnn.cpp +18 -14
  638. data/ext/sources/ggml/src/ggml-zendnn/CMakeLists.txt +1 -1
  639. data/ext/sources/ggml/src/ggml-zendnn/ggml-zendnn.cpp +244 -10
  640. data/ext/sources/ggml/src/ggml.c +146 -42
  641. data/ext/sources/ggml/src/gguf.cpp +173 -28
  642. data/ext/sources/include/parakeet.h +342 -0
  643. data/ext/sources/include/whisper.h +31 -0
  644. data/ext/sources/media/matmul.png +0 -0
  645. data/ext/sources/src/CMakeLists.txt +23 -0
  646. data/ext/sources/src/parakeet-arch.h +188 -0
  647. data/ext/sources/src/parakeet.cpp +3838 -0
  648. data/ext/sources/src/whisper.cpp +220 -26
  649. data/extsources.rb +26 -10
  650. data/lib/whisper/log_settable.rb +33 -0
  651. data/lib/whisper/model/uri.rb +13 -8
  652. data/lib/whisper/output.rb +74 -0
  653. data/sig/whisper.rbs +417 -62
  654. data/test/helper.rb +2 -0
  655. data/test/jfk_reader/jfk_reader.c +50 -7
  656. data/test/test_callback.rb +1 -0
  657. data/test/test_package.rb +6 -5
  658. data/test/test_parakeet.rb +28 -0
  659. data/test/test_parakeet_callback.rb +107 -0
  660. data/test/test_parakeet_context.rb +116 -0
  661. data/test/test_parakeet_context_params.rb +24 -0
  662. data/test/test_parakeet_model.rb +21 -0
  663. data/test/test_parakeet_params.rb +78 -0
  664. data/test/test_parakeet_segment.rb +42 -0
  665. data/test/test_parakeet_token.rb +73 -0
  666. data/test/test_params.rb +2 -0
  667. data/test/test_vad.rb +9 -0
  668. data/test/test_vad_context.rb +2 -2
  669. data/test/test_vad_segment.rb +1 -1
  670. data/test/test_whisper.rb +24 -6
  671. data/whispercpp.gemspec +2 -2
  672. metadata +263 -304
  673. data/ext/sources/bindings/javascript/CMakeLists.txt +0 -41
  674. data/ext/sources/bindings/javascript/emscripten.cpp +0 -93
  675. data/ext/sources/bindings/javascript/libwhisper.worker.js +0 -1
  676. data/ext/sources/bindings/javascript/package.json +0 -26
  677. data/ext/sources/bindings/javascript/whisper.js +0 -19
  678. data/ext/sources/examples/addon.node/CMakeLists.txt +0 -31
  679. data/ext/sources/examples/addon.node/__test__/whisper.spec.js +0 -133
  680. data/ext/sources/examples/addon.node/addon.cpp +0 -557
  681. data/ext/sources/examples/addon.node/index.js +0 -59
  682. data/ext/sources/examples/addon.node/package.json +0 -16
  683. data/ext/sources/examples/addon.node/vad-example.js +0 -132
  684. data/ext/sources/examples/bench.wasm/CMakeLists.txt +0 -49
  685. data/ext/sources/examples/bench.wasm/emscripten.cpp +0 -87
  686. data/ext/sources/examples/bench.wasm/index-tmpl.html +0 -285
  687. data/ext/sources/examples/coi-serviceworker.js +0 -146
  688. data/ext/sources/examples/command/CMakeLists.txt +0 -10
  689. data/ext/sources/examples/command/command.cpp +0 -802
  690. data/ext/sources/examples/command/commands.txt +0 -9
  691. data/ext/sources/examples/command.wasm/CMakeLists.txt +0 -50
  692. data/ext/sources/examples/command.wasm/emscripten.cpp +0 -327
  693. data/ext/sources/examples/command.wasm/index-tmpl.html +0 -415
  694. data/ext/sources/examples/generate-karaoke.sh +0 -57
  695. data/ext/sources/examples/helpers.js +0 -191
  696. data/ext/sources/examples/livestream.sh +0 -112
  697. data/ext/sources/examples/lsp/CMakeLists.txt +0 -10
  698. data/ext/sources/examples/lsp/lsp.cpp +0 -471
  699. data/ext/sources/examples/lsp/whisper.vim +0 -362
  700. data/ext/sources/examples/python/test_whisper_processor.py +0 -7
  701. data/ext/sources/examples/python/whisper_processor.py +0 -54
  702. data/ext/sources/examples/server/bench.js +0 -29
  703. data/ext/sources/examples/server.py +0 -120
  704. data/ext/sources/examples/stream/CMakeLists.txt +0 -10
  705. data/ext/sources/examples/stream/stream.cpp +0 -437
  706. data/ext/sources/examples/stream.wasm/CMakeLists.txt +0 -49
  707. data/ext/sources/examples/stream.wasm/emscripten.cpp +0 -216
  708. data/ext/sources/examples/stream.wasm/index-tmpl.html +0 -491
  709. data/ext/sources/examples/sycl/CMakeLists.txt +0 -9
  710. data/ext/sources/examples/sycl/build.sh +0 -22
  711. data/ext/sources/examples/sycl/ls-sycl-device.cpp +0 -11
  712. data/ext/sources/examples/sycl/run-whisper.sh +0 -17
  713. data/ext/sources/examples/talk-llama/CMakeLists.txt +0 -48
  714. data/ext/sources/examples/talk-llama/eleven-labs.py +0 -80
  715. data/ext/sources/examples/talk-llama/llama-adapter.cpp +0 -488
  716. data/ext/sources/examples/talk-llama/llama-adapter.h +0 -89
  717. data/ext/sources/examples/talk-llama/llama-arch.cpp +0 -2877
  718. data/ext/sources/examples/talk-llama/llama-arch.h +0 -628
  719. data/ext/sources/examples/talk-llama/llama-batch.cpp +0 -919
  720. data/ext/sources/examples/talk-llama/llama-batch.h +0 -173
  721. data/ext/sources/examples/talk-llama/llama-chat.cpp +0 -896
  722. data/ext/sources/examples/talk-llama/llama-chat.h +0 -71
  723. data/ext/sources/examples/talk-llama/llama-context.cpp +0 -3633
  724. data/ext/sources/examples/talk-llama/llama-context.h +0 -359
  725. data/ext/sources/examples/talk-llama/llama-cparams.cpp +0 -5
  726. data/ext/sources/examples/talk-llama/llama-cparams.h +0 -47
  727. data/ext/sources/examples/talk-llama/llama-ext.h +0 -12
  728. data/ext/sources/examples/talk-llama/llama-grammar.cpp +0 -1464
  729. data/ext/sources/examples/talk-llama/llama-grammar.h +0 -194
  730. data/ext/sources/examples/talk-llama/llama-graph.cpp +0 -2735
  731. data/ext/sources/examples/talk-llama/llama-graph.h +0 -1031
  732. data/ext/sources/examples/talk-llama/llama-hparams.cpp +0 -258
  733. data/ext/sources/examples/talk-llama/llama-hparams.h +0 -353
  734. data/ext/sources/examples/talk-llama/llama-impl.cpp +0 -171
  735. data/ext/sources/examples/talk-llama/llama-impl.h +0 -75
  736. data/ext/sources/examples/talk-llama/llama-io.cpp +0 -15
  737. data/ext/sources/examples/talk-llama/llama-io.h +0 -35
  738. data/ext/sources/examples/talk-llama/llama-kv-cache-iswa.cpp +0 -330
  739. data/ext/sources/examples/talk-llama/llama-kv-cache-iswa.h +0 -137
  740. data/ext/sources/examples/talk-llama/llama-kv-cache.cpp +0 -2285
  741. data/ext/sources/examples/talk-llama/llama-kv-cache.h +0 -389
  742. data/ext/sources/examples/talk-llama/llama-kv-cells.h +0 -533
  743. data/ext/sources/examples/talk-llama/llama-memory-hybrid-iswa.cpp +0 -275
  744. data/ext/sources/examples/talk-llama/llama-memory-hybrid-iswa.h +0 -140
  745. data/ext/sources/examples/talk-llama/llama-memory-hybrid.cpp +0 -268
  746. data/ext/sources/examples/talk-llama/llama-memory-hybrid.h +0 -139
  747. data/ext/sources/examples/talk-llama/llama-memory-recurrent.cpp +0 -1165
  748. data/ext/sources/examples/talk-llama/llama-memory-recurrent.h +0 -182
  749. data/ext/sources/examples/talk-llama/llama-memory.cpp +0 -59
  750. data/ext/sources/examples/talk-llama/llama-memory.h +0 -122
  751. data/ext/sources/examples/talk-llama/llama-mmap.cpp +0 -752
  752. data/ext/sources/examples/talk-llama/llama-mmap.h +0 -73
  753. data/ext/sources/examples/talk-llama/llama-model-loader.cpp +0 -1655
  754. data/ext/sources/examples/talk-llama/llama-model-loader.h +0 -206
  755. data/ext/sources/examples/talk-llama/llama-model-saver.cpp +0 -299
  756. data/ext/sources/examples/talk-llama/llama-model-saver.h +0 -40
  757. data/ext/sources/examples/talk-llama/llama-model.cpp +0 -9056
  758. data/ext/sources/examples/talk-llama/llama-model.h +0 -597
  759. data/ext/sources/examples/talk-llama/llama-quant.cpp +0 -1304
  760. data/ext/sources/examples/talk-llama/llama-quant.h +0 -1
  761. data/ext/sources/examples/talk-llama/llama-sampler.cpp +0 -3885
  762. data/ext/sources/examples/talk-llama/llama-sampler.h +0 -42
  763. data/ext/sources/examples/talk-llama/llama-vocab.cpp +0 -3970
  764. data/ext/sources/examples/talk-llama/llama-vocab.h +0 -187
  765. data/ext/sources/examples/talk-llama/llama.cpp +0 -1194
  766. data/ext/sources/examples/talk-llama/llama.h +0 -1573
  767. data/ext/sources/examples/talk-llama/models/afmoe.cpp +0 -190
  768. data/ext/sources/examples/talk-llama/models/apertus.cpp +0 -125
  769. data/ext/sources/examples/talk-llama/models/arcee.cpp +0 -135
  770. data/ext/sources/examples/talk-llama/models/arctic.cpp +0 -137
  771. data/ext/sources/examples/talk-llama/models/arwkv7.cpp +0 -86
  772. data/ext/sources/examples/talk-llama/models/baichuan.cpp +0 -123
  773. data/ext/sources/examples/talk-llama/models/bailingmoe.cpp +0 -143
  774. data/ext/sources/examples/talk-llama/models/bailingmoe2.cpp +0 -133
  775. data/ext/sources/examples/talk-llama/models/bert.cpp +0 -184
  776. data/ext/sources/examples/talk-llama/models/bitnet.cpp +0 -145
  777. data/ext/sources/examples/talk-llama/models/bloom.cpp +0 -101
  778. data/ext/sources/examples/talk-llama/models/chameleon.cpp +0 -178
  779. data/ext/sources/examples/talk-llama/models/chatglm.cpp +0 -132
  780. data/ext/sources/examples/talk-llama/models/codeshell.cpp +0 -111
  781. data/ext/sources/examples/talk-llama/models/cogvlm.cpp +0 -102
  782. data/ext/sources/examples/talk-llama/models/cohere2-iswa.cpp +0 -134
  783. data/ext/sources/examples/talk-llama/models/command-r.cpp +0 -122
  784. data/ext/sources/examples/talk-llama/models/dbrx.cpp +0 -122
  785. data/ext/sources/examples/talk-llama/models/deci.cpp +0 -135
  786. data/ext/sources/examples/talk-llama/models/deepseek.cpp +0 -142
  787. data/ext/sources/examples/talk-llama/models/deepseek2.cpp +0 -262
  788. data/ext/sources/examples/talk-llama/models/delta-net-base.cpp +0 -445
  789. data/ext/sources/examples/talk-llama/models/dots1.cpp +0 -132
  790. data/ext/sources/examples/talk-llama/models/dream.cpp +0 -105
  791. data/ext/sources/examples/talk-llama/models/ernie4-5-moe.cpp +0 -148
  792. data/ext/sources/examples/talk-llama/models/ernie4-5.cpp +0 -110
  793. data/ext/sources/examples/talk-llama/models/eurobert.cpp +0 -97
  794. data/ext/sources/examples/talk-llama/models/exaone-moe.cpp +0 -145
  795. data/ext/sources/examples/talk-llama/models/exaone.cpp +0 -114
  796. data/ext/sources/examples/talk-llama/models/exaone4.cpp +0 -123
  797. data/ext/sources/examples/talk-llama/models/falcon-h1.cpp +0 -111
  798. data/ext/sources/examples/talk-llama/models/falcon.cpp +0 -120
  799. data/ext/sources/examples/talk-llama/models/gemma-embedding.cpp +0 -116
  800. data/ext/sources/examples/talk-llama/models/gemma.cpp +0 -112
  801. data/ext/sources/examples/talk-llama/models/gemma2-iswa.cpp +0 -128
  802. data/ext/sources/examples/talk-llama/models/gemma3.cpp +0 -155
  803. data/ext/sources/examples/talk-llama/models/gemma3n-iswa.cpp +0 -384
  804. data/ext/sources/examples/talk-llama/models/glm4-moe.cpp +0 -170
  805. data/ext/sources/examples/talk-llama/models/glm4.cpp +0 -157
  806. data/ext/sources/examples/talk-llama/models/gpt2.cpp +0 -105
  807. data/ext/sources/examples/talk-llama/models/gptneox.cpp +0 -144
  808. data/ext/sources/examples/talk-llama/models/granite-hybrid.cpp +0 -195
  809. data/ext/sources/examples/talk-llama/models/granite.cpp +0 -210
  810. data/ext/sources/examples/talk-llama/models/grok.cpp +0 -159
  811. data/ext/sources/examples/talk-llama/models/grovemoe.cpp +0 -139
  812. data/ext/sources/examples/talk-llama/models/hunyuan-dense.cpp +0 -132
  813. data/ext/sources/examples/talk-llama/models/hunyuan-moe.cpp +0 -153
  814. data/ext/sources/examples/talk-llama/models/internlm2.cpp +0 -120
  815. data/ext/sources/examples/talk-llama/models/jais.cpp +0 -86
  816. data/ext/sources/examples/talk-llama/models/jais2.cpp +0 -123
  817. data/ext/sources/examples/talk-llama/models/jamba.cpp +0 -106
  818. data/ext/sources/examples/talk-llama/models/kimi-linear.cpp +0 -381
  819. data/ext/sources/examples/talk-llama/models/lfm2.cpp +0 -196
  820. data/ext/sources/examples/talk-llama/models/llada-moe.cpp +0 -122
  821. data/ext/sources/examples/talk-llama/models/llada.cpp +0 -99
  822. data/ext/sources/examples/talk-llama/models/llama-iswa.cpp +0 -178
  823. data/ext/sources/examples/talk-llama/models/llama.cpp +0 -175
  824. data/ext/sources/examples/talk-llama/models/maincoder.cpp +0 -117
  825. data/ext/sources/examples/talk-llama/models/mamba-base.cpp +0 -289
  826. data/ext/sources/examples/talk-llama/models/mamba.cpp +0 -54
  827. data/ext/sources/examples/talk-llama/models/mimo2-iswa.cpp +0 -129
  828. data/ext/sources/examples/talk-llama/models/minicpm3.cpp +0 -200
  829. data/ext/sources/examples/talk-llama/models/minimax-m2.cpp +0 -123
  830. data/ext/sources/examples/talk-llama/models/mistral3.cpp +0 -160
  831. data/ext/sources/examples/talk-llama/models/models.h +0 -704
  832. data/ext/sources/examples/talk-llama/models/modern-bert.cpp +0 -109
  833. data/ext/sources/examples/talk-llama/models/mpt.cpp +0 -126
  834. data/ext/sources/examples/talk-llama/models/nemotron-h.cpp +0 -162
  835. data/ext/sources/examples/talk-llama/models/nemotron.cpp +0 -122
  836. data/ext/sources/examples/talk-llama/models/neo-bert.cpp +0 -104
  837. data/ext/sources/examples/talk-llama/models/olmo.cpp +0 -121
  838. data/ext/sources/examples/talk-llama/models/olmo2.cpp +0 -150
  839. data/ext/sources/examples/talk-llama/models/olmoe.cpp +0 -124
  840. data/ext/sources/examples/talk-llama/models/openai-moe-iswa.cpp +0 -127
  841. data/ext/sources/examples/talk-llama/models/openelm.cpp +0 -124
  842. data/ext/sources/examples/talk-llama/models/orion.cpp +0 -123
  843. data/ext/sources/examples/talk-llama/models/paddleocr.cpp +0 -122
  844. data/ext/sources/examples/talk-llama/models/pangu-embedded.cpp +0 -121
  845. data/ext/sources/examples/talk-llama/models/phi2.cpp +0 -121
  846. data/ext/sources/examples/talk-llama/models/phi3.cpp +0 -152
  847. data/ext/sources/examples/talk-llama/models/plamo.cpp +0 -110
  848. data/ext/sources/examples/talk-llama/models/plamo2.cpp +0 -320
  849. data/ext/sources/examples/talk-llama/models/plamo3.cpp +0 -128
  850. data/ext/sources/examples/talk-llama/models/plm.cpp +0 -169
  851. data/ext/sources/examples/talk-llama/models/qwen.cpp +0 -108
  852. data/ext/sources/examples/talk-llama/models/qwen2.cpp +0 -126
  853. data/ext/sources/examples/talk-llama/models/qwen2moe.cpp +0 -151
  854. data/ext/sources/examples/talk-llama/models/qwen2vl.cpp +0 -117
  855. data/ext/sources/examples/talk-llama/models/qwen3.cpp +0 -120
  856. data/ext/sources/examples/talk-llama/models/qwen35.cpp +0 -381
  857. data/ext/sources/examples/talk-llama/models/qwen35moe.cpp +0 -422
  858. data/ext/sources/examples/talk-llama/models/qwen3moe.cpp +0 -131
  859. data/ext/sources/examples/talk-llama/models/qwen3next.cpp +0 -525
  860. data/ext/sources/examples/talk-llama/models/qwen3vl-moe.cpp +0 -140
  861. data/ext/sources/examples/talk-llama/models/qwen3vl.cpp +0 -132
  862. data/ext/sources/examples/talk-llama/models/refact.cpp +0 -94
  863. data/ext/sources/examples/talk-llama/models/rnd1.cpp +0 -126
  864. data/ext/sources/examples/talk-llama/models/rwkv6-base.cpp +0 -164
  865. data/ext/sources/examples/talk-llama/models/rwkv6.cpp +0 -94
  866. data/ext/sources/examples/talk-llama/models/rwkv6qwen2.cpp +0 -86
  867. data/ext/sources/examples/talk-llama/models/rwkv7-base.cpp +0 -137
  868. data/ext/sources/examples/talk-llama/models/rwkv7.cpp +0 -90
  869. data/ext/sources/examples/talk-llama/models/seed-oss.cpp +0 -124
  870. data/ext/sources/examples/talk-llama/models/smallthinker.cpp +0 -126
  871. data/ext/sources/examples/talk-llama/models/smollm3.cpp +0 -128
  872. data/ext/sources/examples/talk-llama/models/stablelm.cpp +0 -146
  873. data/ext/sources/examples/talk-llama/models/starcoder.cpp +0 -100
  874. data/ext/sources/examples/talk-llama/models/starcoder2.cpp +0 -121
  875. data/ext/sources/examples/talk-llama/models/step35-iswa.cpp +0 -165
  876. data/ext/sources/examples/talk-llama/models/t5-dec.cpp +0 -166
  877. data/ext/sources/examples/talk-llama/models/t5-enc.cpp +0 -96
  878. data/ext/sources/examples/talk-llama/models/wavtokenizer-dec.cpp +0 -149
  879. data/ext/sources/examples/talk-llama/models/xverse.cpp +0 -108
  880. data/ext/sources/examples/talk-llama/prompts/talk-alpaca.txt +0 -23
  881. data/ext/sources/examples/talk-llama/speak +0 -40
  882. data/ext/sources/examples/talk-llama/speak.bat +0 -1
  883. data/ext/sources/examples/talk-llama/speak.ps1 +0 -14
  884. data/ext/sources/examples/talk-llama/talk-llama.cpp +0 -813
  885. data/ext/sources/examples/talk-llama/unicode-data.cpp +0 -7034
  886. data/ext/sources/examples/talk-llama/unicode-data.h +0 -20
  887. data/ext/sources/examples/talk-llama/unicode.cpp +0 -1103
  888. data/ext/sources/examples/talk-llama/unicode.h +0 -111
  889. data/ext/sources/examples/wchess/CMakeLists.txt +0 -10
  890. data/ext/sources/examples/wchess/libwchess/CMakeLists.txt +0 -19
  891. data/ext/sources/examples/wchess/libwchess/Chessboard.cpp +0 -803
  892. data/ext/sources/examples/wchess/libwchess/Chessboard.h +0 -33
  893. data/ext/sources/examples/wchess/libwchess/WChess.cpp +0 -193
  894. data/ext/sources/examples/wchess/libwchess/WChess.h +0 -63
  895. data/ext/sources/examples/wchess/libwchess/test-chessboard.cpp +0 -117
  896. data/ext/sources/examples/wchess/wchess.cmd/CMakeLists.txt +0 -8
  897. data/ext/sources/examples/wchess/wchess.cmd/wchess.cmd.cpp +0 -253
  898. data/ext/sources/examples/whisper.wasm/CMakeLists.txt +0 -50
  899. data/ext/sources/examples/whisper.wasm/emscripten.cpp +0 -118
  900. data/ext/sources/examples/whisper.wasm/index-tmpl.html +0 -659
  901. data/ext/sources/ggml/src/ggml-cuda/template-instances/generate_cu_files.py +0 -99
  902. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-msg.h +0 -155
  903. data/ext/sources/ggml/src/ggml-hexagon/op-desc.h +0 -153
  904. data/ext/sources/ggml/src/ggml-opencl/kernels/embed_kernel.py +0 -26
  905. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/eliminate_zp.cpp +0 -123
  906. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/eliminate_zp.h +0 -17
  907. data/ext/sources/ggml/src/ggml-virtgpu/regenerate_remoting.py +0 -333
  908. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/abs.comp +0 -21
  909. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ceil.comp +0 -22
  910. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/clamp.comp +0 -17
  911. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/cos.comp +0 -17
  912. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/elu.comp +0 -27
  913. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/exp.comp +0 -21
  914. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/floor.comp +0 -22
  915. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu.comp +0 -25
  916. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_erf.comp +0 -39
  917. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_quick.comp +0 -23
  918. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardsigmoid.comp +0 -22
  919. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardswish.comp +0 -22
  920. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/leaky_relu.comp +0 -22
  921. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/neg.comp +0 -20
  922. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/relu.comp +0 -21
  923. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/round.comp +0 -29
  924. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rte.glsl +0 -5
  925. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sgn.comp +0 -21
  926. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sigmoid.comp +0 -20
  927. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/silu.comp +0 -22
  928. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sin.comp +0 -17
  929. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/softplus.comp +0 -23
  930. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sqrt.comp +0 -17
  931. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/square.comp +0 -17
  932. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/step.comp +0 -22
  933. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tanh.comp +0 -20
  934. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/trunc.comp +0 -22
  935. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/xielu.comp +0 -35
  936. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/embed_wgsl.py +0 -182
  937. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/glu.tmpl.wgsl +0 -323
  938. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat.wgsl +0 -718
  939. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm.wgsl +0 -123
  940. data/ext/sources/tests/CMakeLists.txt +0 -112
  941. data/ext/sources/tests/earnings21/eval.mk +0 -58
  942. data/ext/sources/tests/earnings21/eval.py +0 -68
  943. data/ext/sources/tests/earnings21/normalizers/__init__.py +0 -2
  944. data/ext/sources/tests/earnings21/normalizers/basic.py +0 -80
  945. data/ext/sources/tests/earnings21/normalizers/english.json +0 -1741
  946. data/ext/sources/tests/earnings21/normalizers/english.py +0 -550
  947. data/ext/sources/tests/earnings21/requirements.txt +0 -6
  948. data/ext/sources/tests/en-0-ref.txt +0 -1
  949. data/ext/sources/tests/en-1-ref.txt +0 -1
  950. data/ext/sources/tests/en-2-ref.txt +0 -1
  951. data/ext/sources/tests/es-0-ref.txt +0 -1
  952. data/ext/sources/tests/librispeech/eval.mk +0 -39
  953. data/ext/sources/tests/librispeech/eval.py +0 -47
  954. data/ext/sources/tests/librispeech/normalizers/__init__.py +0 -2
  955. data/ext/sources/tests/librispeech/normalizers/basic.py +0 -80
  956. data/ext/sources/tests/librispeech/normalizers/english.json +0 -1741
  957. data/ext/sources/tests/librispeech/normalizers/english.py +0 -550
  958. data/ext/sources/tests/librispeech/requirements.txt +0 -6
  959. data/ext/sources/tests/run-tests.sh +0 -130
  960. data/ext/sources/tests/test-c.c +0 -3
  961. data/ext/sources/tests/test-vad-full.cpp +0 -56
  962. data/ext/sources/tests/test-vad.cpp +0 -83
  963. data/ext/sources/tests/test-whisper.js +0 -58
  964. data/lib/whisper/context.rb +0 -15
  965. data/lib/whisper/segment.rb +0 -58
@@ -19,6 +19,7 @@
19
19
  #include <cstdlib>
20
20
  #include <float.h>
21
21
  #include <limits>
22
+ #include <optional>
22
23
  #include <stdint.h>
23
24
  #include <stdio.h>
24
25
  #include <vector>
@@ -30,9 +31,18 @@
30
31
  #include <regex>
31
32
 
32
33
  #include <sycl/sycl.hpp>
34
+ #include <sycl/backend.hpp>
35
+ #ifdef GGML_SYCL_SUPPORT_LEVEL_ZERO_API
36
+ #include <level_zero/ze_api.h>
37
+ #endif
33
38
  #if defined(GGML_SYCL_GRAPH) && SYCL_EXT_ONEAPI_ASYNC_MEMORY_ALLOC
34
39
  # include <sycl/ext/oneapi/experimental/async_alloc/async_alloc.hpp>
35
40
  #endif
41
+ #if SYCL_EXT_ONEAPI_VIRTUAL_MEM
42
+ # include <sycl/ext/oneapi/virtual_mem/physical_mem.hpp>
43
+ # include <sycl/ext/oneapi/virtual_mem/virtual_mem.hpp>
44
+ # define GGML_SYCL_SUPPORT_VMM
45
+ #endif
36
46
  #include <sycl/half_type.hpp>
37
47
 
38
48
  #include "ggml.h"
@@ -44,7 +54,6 @@
44
54
  #include "ggml-sycl/backend.hpp"
45
55
  #include "ggml-sycl/common.hpp"
46
56
  #include "ggml-sycl/element_wise.hpp"
47
- #include "ggml-sycl/gated_delta_net.hpp"
48
57
  #include "ggml-sycl/gemm.hpp"
49
58
  #include "ggml-sycl/getrows.hpp"
50
59
  #include "ggml-sycl/norm.hpp"
@@ -53,19 +62,36 @@
53
62
  #include "ggml-sycl/repeat_back.hpp"
54
63
  #include "ggml-sycl/set_rows.hpp"
55
64
  #include "ggml-sycl/set.hpp"
65
+ #include "ggml-sycl/conv2d.hpp"
66
+ #include "ggml-sycl/conv2d-dw.hpp"
67
+ #include "ggml-sycl/conv2d-transpose.hpp"
56
68
  #include "ggml-sycl/ssm_conv.hpp"
57
69
  #include "ggml-sycl/sycl_hw.hpp"
70
+ #include "ggml-sycl/ssm_scan.hpp"
71
+ #include "ggml-sycl/fill.hpp"
72
+ #include "ggml-sycl/cumsum.hpp"
73
+ #include "ggml-sycl/diag.hpp"
74
+ #include "ggml-sycl/solve_tri.hpp"
75
+ #include "ggml-sycl/gated_delta_net.hpp"
76
+ #include "ggml-sycl/pool.hpp"
77
+ #include "ggml-sycl/cross_entropy_loss.hpp"
58
78
 
79
+ #define MEM_SIZE_2M 0x00200000
80
+ #define MEM_SIZE_1G 0x40000000
59
81
 
60
82
  static bool g_sycl_loaded = false;
61
83
  int g_ggml_sycl_debug = 0;
62
- int g_ggml_sycl_disable_optimize = 0;
63
- int g_ggml_sycl_disable_graph = 0;
64
- int g_ggml_sycl_disable_dnn = 0;
84
+ int g_ggml_sycl_enable_optimize = 1;
85
+ int g_ggml_sycl_enable_graph = 0;
86
+ int g_ggml_sycl_enable_dnn = 1;
87
+ int g_ggml_sycl_enable_vmm = 1;
65
88
  int g_ggml_sycl_prioritize_dmmv = 0;
66
89
  int g_ggml_sycl_use_async_mem_op = 0;
90
+ int g_ggml_sycl_use_async_mem_op_requested = 1;
91
+ int g_ggml_sycl_use_level_zero_api = 0;
67
92
  int g_ggml_sycl_enable_flash_attention = 1;
68
-
93
+ int g_ggml_sycl_dev2dev_memcpy = DEV2DEV_MEMCPY_SYCL;
94
+ int g_ggml_sycl_usm_system = 0;
69
95
 
70
96
  static ggml_sycl_device_info ggml_sycl_init() {
71
97
  ggml_sycl_device_info info = {};
@@ -86,13 +112,30 @@ static ggml_sycl_device_info ggml_sycl_init() {
86
112
  // GGML_LOG_INFO("%s: SYCL_USE_XMX: no\n", __func__);
87
113
  // #endif
88
114
  for (int i = 0; i < info.device_count; ++i) {
89
- info.devices[i].vmm = 0;
90
115
  dpct::device_info prop;
91
- sycl::device device = dpct::dev_mgr::instance().get_device(i);
116
+ auto & device = dpct::dev_mgr::instance().get_device(i);
92
117
 
93
118
  SYCL_CHECK(CHECK_TRY_ERROR(dpct::get_device_info(
94
119
  prop, device)));
95
120
 
121
+ #if !defined(GGML_SYCL_SUPPORT_VMM)
122
+ info.devices[i].vmm = 0;
123
+ #else
124
+ info.devices[i].vmm = device.has(sycl::aspect::ext_oneapi_virtual_mem);
125
+ if (info.devices[i].vmm) {
126
+ // NB: SYCL's get_mem_granularity always returns the _minimum_ granularity,
127
+ // but the L0 API requires a larger page size for allocs above 2 MiB and
128
+ // rejects non-multiples with UR_RESULT_ERROR_INVALID_VALUE [sic].
129
+ // Here we clamp it to 2 MiB for simplicity, but other devices may require
130
+ // calling zeVirtualMemQueryPageSize or yet unexposed public API.
131
+ const size_t physical_page = 2ull << 20; // 2 MiB
132
+ info.devices[i].vmm_granularity = std::max<size_t>(
133
+ sycl::ext::oneapi::experimental::get_mem_granularity(
134
+ device, sycl::context(device)),
135
+ physical_page);
136
+ }
137
+ #endif
138
+
96
139
  info.default_tensor_split[i] = total_vram;
97
140
  total_vram += prop.get_global_mem_size();
98
141
 
@@ -102,15 +145,43 @@ static ggml_sycl_device_info ggml_sycl_init() {
102
145
  info.devices[i].opt_feature.reorder = device.ext_oneapi_architecture_is(syclex::arch_category::intel_gpu);
103
146
  info.devices[i].smpbo = prop.get_local_mem_size();
104
147
  info.devices[i].warp_size = WARP_SIZE;
148
+ info.devices[i].usm_system_support = device.has(sycl::aspect::usm_system_allocations);
105
149
 
106
150
  info.max_work_group_sizes[i] = prop.get_max_work_group_size();
107
151
  info.devices[i].max_wg_per_cu = info.max_work_group_sizes[i] / prop.get_max_compute_units();
152
+ info.devices[i].hw_info = get_device_hw_info(&device);
153
+
154
+ // Only check GPU devices; CPU devices use OpenCL and would otherwise
155
+ // disable Level Zero for the GPUs on systems without ONEAPI_DEVICE_SELECTOR set.
156
+ if (device.is_gpu() && device.default_queue().get_backend() != sycl::backend::ext_oneapi_level_zero) {
157
+ GGML_LOG_WARN("SYCL GPU device %d does not use Level Zero backend, disabling Level Zero memory API\n", i);
158
+ info.ext_oneapi_level_zero = false;
159
+ }
108
160
 
161
+ #ifdef GGML_SYCL_SUPPORT_LEVEL_ZERO_API
162
+ if (info.ext_oneapi_level_zero && device.is_gpu() && device.default_queue().get_backend() == sycl::backend::ext_oneapi_level_zero) {
163
+ ze_device_handle_t ze_dev = sycl::get_native<sycl::backend::ext_oneapi_level_zero>(device.default_queue().get_device());
164
+ ze_device_properties_t props = {};
165
+ props.stype = ZE_STRUCTURE_TYPE_DEVICE_PROPERTIES;
166
+ ze_result_t r = zeDeviceGetProperties(ze_dev, &props);
167
+ info.devices[i].l0_discrete_gpu = r == ZE_RESULT_SUCCESS && !(props.flags & ZE_DEVICE_PROPERTY_FLAG_INTEGRATED);
168
+ }
169
+ #endif
109
170
  }
110
171
 
111
172
  for (int id = 0; id < info.device_count; ++id) {
112
173
  info.default_tensor_split[id] /= total_vram;
113
174
  }
175
+
176
+ #ifdef GGML_SYCL_SUPPORT_LEVEL_ZERO_API
177
+ // Large buffers can be allocated before ggml_check_sycl() initializes other
178
+ // g_ggml_sycl_enable_* globals, so initialize this one as early as we can.
179
+ g_ggml_sycl_use_level_zero_api =
180
+ info.ext_oneapi_level_zero && ggml_sycl_get_env("GGML_SYCL_USE_LEVEL_ZERO_API", 1);
181
+ #else
182
+ g_ggml_sycl_use_level_zero_api = 0;
183
+ #endif
184
+
114
185
  return info;
115
186
  }
116
187
 
@@ -195,74 +266,93 @@ void ggml_backend_sycl_print_sycl_devices() {
195
266
  print_device_opt_feature(device_count);
196
267
  }
197
268
 
198
- static inline int get_sycl_env(const char *env_name, int default_val) {
199
- char *user_device_string = getenv(env_name);
200
- int user_number = default_val;
201
-
202
- unsigned n;
203
- if (user_device_string != NULL &&
204
- sscanf(user_device_string, " %u", &n) == 1) {
205
- user_number = (int)n;
269
+ static const char* dev2dev_int2str(int dev2dev) {
270
+ if (dev2dev == DEV2DEV_MEMCPY_SYCL) {
271
+ return "SYCL API";
272
+ } else if (dev2dev == DEV2DEV_MEMCPY_L0) {
273
+ return "Level Zero API";
206
274
  } else {
207
- user_number = default_val;
275
+ return "Unknown";
208
276
  }
209
- return user_number;
210
277
  }
211
278
 
212
279
  static void ggml_check_sycl() try {
213
280
  static bool initialized = false;
214
281
 
215
282
  if (!initialized) {
216
- g_ggml_sycl_debug = get_sycl_env("GGML_SYCL_DEBUG", 0);
217
- g_ggml_sycl_disable_optimize = get_sycl_env("GGML_SYCL_DISABLE_OPT", 0);
218
- g_ggml_sycl_disable_graph = get_sycl_env("GGML_SYCL_DISABLE_GRAPH", 1);
219
- g_ggml_sycl_disable_dnn = get_sycl_env("GGML_SYCL_DISABLE_DNN", 0);
220
- g_ggml_sycl_prioritize_dmmv = get_sycl_env("GGML_SYCL_PRIORITIZE_DMMV", 0);
283
+ g_ggml_sycl_debug = ggml_sycl_get_env("GGML_SYCL_DEBUG", 0);
284
+ g_ggml_sycl_enable_optimize = ggml_sycl_get_env("GGML_SYCL_ENABLE_OPT", 1);
285
+ g_ggml_sycl_enable_graph = ggml_sycl_get_env("GGML_SYCL_ENABLE_GRAPH", 0);
286
+ g_ggml_sycl_enable_dnn = ggml_sycl_get_env("GGML_SYCL_ENABLE_DNN", 1);
287
+ g_ggml_sycl_enable_vmm = ggml_sycl_get_env("GGML_SYCL_ENABLE_VMM", 1);
288
+ g_ggml_sycl_prioritize_dmmv = ggml_sycl_get_env("GGML_SYCL_PRIORITIZE_DMMV", 0);
289
+
290
+ g_ggml_sycl_dev2dev_memcpy = ggml_sycl_get_env("GGML_SYCL_DEV2DEV_MEMCPY", DEV2DEV_MEMCPY_SYCL);
291
+ if (g_ggml_sycl_use_level_zero_api == 0) {
292
+ g_ggml_sycl_dev2dev_memcpy = DEV2DEV_MEMCPY_SYCL;
293
+ }
221
294
 
222
295
  #ifdef SYCL_FLASH_ATTN
223
- g_ggml_sycl_enable_flash_attention = get_sycl_env("GGML_SYCL_ENABLE_FLASH_ATTN", 1);
296
+ g_ggml_sycl_enable_flash_attention = ggml_sycl_get_env("GGML_SYCL_ENABLE_FLASH_ATTN", 1);
224
297
  #else
225
298
  g_ggml_sycl_enable_flash_attention = 0;
226
299
  #endif
227
300
 
301
+ g_ggml_sycl_usm_system = ggml_sycl_get_env("GGML_SYCL_USM_SYSTEM", 0);
302
+
228
303
  GGML_SYCL_DEBUG("[SYCL] call ggml_check_sycl\n");
229
304
 
230
305
  GGML_LOG_INFO("Build with Macros:\n");
231
- #if defined(GGML_SYCL_FORCE_MMQ)
232
- GGML_LOG_INFO(" GGML_SYCL_FORCE_MMQ: yes\n");
306
+ #if defined(GGML_SYCL_DNNL)
307
+ GGML_LOG_INFO(" GGML_SYCL_DNNL: yes\n");
233
308
  #else
234
- GGML_LOG_INFO(" GGML_SYCL_FORCE_MMQ: no\n");
309
+ GGML_LOG_INFO(" GGML_SYCL_DNNL: no\n");
235
310
  #endif
311
+
236
312
  #if defined(GGML_SYCL_F16)
237
313
  GGML_LOG_INFO(" GGML_SYCL_F16: yes\n");
238
314
  #else
239
315
  GGML_LOG_INFO(" GGML_SYCL_F16: no\n");
240
316
  #endif
317
+
318
+ #if defined(GGML_SYCL_FORCE_MMQ)
319
+ GGML_LOG_INFO(" GGML_SYCL_FORCE_MMQ: yes\n");
320
+ #else
321
+ GGML_LOG_INFO(" GGML_SYCL_FORCE_MMQ: no\n");
322
+ #endif
323
+
241
324
  #if defined(GGML_SYCL_GRAPH)
242
325
  GGML_LOG_INFO(" GGML_SYCL_GRAPH: yes\n");
243
326
  #else
244
327
  GGML_LOG_INFO(" GGML_SYCL_GRAPH: no\n");
245
328
  #endif
246
- #if defined(GGML_SYCL_DNNL)
247
- GGML_LOG_INFO(" GGML_SYCL_DNNL: yes\n");
329
+
330
+ #if defined(GGML_SYCL_SUPPORT_LEVEL_ZERO_API)
331
+ GGML_LOG_INFO(" GGML_SYCL_SUPPORT_LEVEL_ZERO_API: yes\n");
248
332
  #else
249
- GGML_LOG_INFO(" GGML_SYCL_DNNL: no\n");
333
+ GGML_LOG_INFO(" GGML_SYCL_SUPPORT_LEVEL_ZERO_API: no\n");
334
+ #endif
335
+ #if defined(GGML_SYCL_SUPPORT_VMM)
336
+ GGML_LOG_INFO(" GGML_SYCL_SUPPORT_VMM: yes\n");
337
+ #else
338
+ GGML_LOG_INFO(" GGML_SYCL_SUPPORT_VMM: no\n");
250
339
  #endif
251
340
 
252
341
  GGML_LOG_INFO("Running with Environment Variables:\n");
253
342
  GGML_LOG_INFO(" GGML_SYCL_DEBUG: %d\n", g_ggml_sycl_debug);
254
- GGML_LOG_INFO(" GGML_SYCL_DISABLE_OPT: %d\n", g_ggml_sycl_disable_optimize);
255
- #ifdef GGML_SYCL_GRAPH
256
- GGML_LOG_INFO(" GGML_SYCL_DISABLE_GRAPH: %d\n", g_ggml_sycl_disable_graph);
343
+
344
+ #ifdef GGML_SYCL_SUPPORT_LEVEL_ZERO_API
345
+ GGML_LOG_INFO(" GGML_SYCL_DEV2DEV_MEMCPY: %d (%s)\n", g_ggml_sycl_dev2dev_memcpy, dev2dev_int2str(g_ggml_sycl_dev2dev_memcpy));
257
346
  #else
258
- GGML_LOG_INFO(" GGML_SYCL_DISABLE_GRAPH: graph disabled by compile flag\n");
347
+ GGML_LOG_INFO(" GGML_SYCL_DEV2DEV_MEMCPY: %d (%s), enable to SYCL API since missing GGML_SYCL_SUPPORT_LEVEL_ZERO_API\n",
348
+ g_ggml_sycl_dev2dev_memcpy, dev2dev_int2str(g_ggml_sycl_dev2dev_memcpy));
259
349
  #endif
260
- #if GGML_SYCL_DNNL
261
- GGML_LOG_INFO(" GGML_SYCL_DISABLE_DNN: %d\n", g_ggml_sycl_disable_dnn);
350
+
351
+ #if defined(GGML_SYCL_DNNL)
352
+ GGML_LOG_INFO(" GGML_SYCL_ENABLE_DNN: %d\n", g_ggml_sycl_enable_dnn);
262
353
  #else
263
- GGML_LOG_INFO(" GGML_SYCL_DISABLE_DNN: DNN disabled by compile flag\n");
354
+ GGML_LOG_INFO(" GGML_SYCL_ENABLE_DNN: DNN disabled by compile flag\n");
264
355
  #endif
265
- GGML_LOG_INFO(" GGML_SYCL_PRIORITIZE_DMMV: %d\n", g_ggml_sycl_prioritize_dmmv);
266
356
 
267
357
  #ifdef SYCL_FLASH_ATTN
268
358
  GGML_LOG_INFO(" GGML_SYCL_ENABLE_FLASH_ATTN: %d\n", g_ggml_sycl_enable_flash_attention);
@@ -271,6 +361,33 @@ static void ggml_check_sycl() try {
271
361
  g_ggml_sycl_enable_flash_attention);
272
362
  #endif
273
363
 
364
+ #ifdef GGML_SYCL_GRAPH
365
+ GGML_LOG_INFO(" GGML_SYCL_ENABLE_GRAPH: %d\n", g_ggml_sycl_enable_graph);
366
+ #else
367
+ GGML_LOG_INFO(" GGML_SYCL_ENABLE_GRAPH: graph disabled by compile flag\n");
368
+ #endif
369
+
370
+ GGML_LOG_INFO(" GGML_SYCL_ENABLE_OPT: %d\n", g_ggml_sycl_enable_optimize);
371
+
372
+ #if defined(GGML_SYCL_SUPPORT_VMM)
373
+ GGML_LOG_INFO(" GGML_SYCL_ENABLE_VMM: %d\n", g_ggml_sycl_enable_vmm);
374
+ #else
375
+ GGML_LOG_INFO(" GGML_SYCL_ENABLE_VMM: virtual memory extension is not available\n");
376
+ #endif
377
+
378
+ GGML_LOG_INFO(" GGML_SYCL_PRIORITIZE_DMMV: %d\n", g_ggml_sycl_prioritize_dmmv);
379
+
380
+ g_ggml_sycl_use_async_mem_op_requested = ggml_sycl_get_env("GGML_SYCL_USE_ASYNC_MEM_OP", 1);
381
+ GGML_LOG_INFO(" GGML_SYCL_USE_ASYNC_MEM_OP: %d\n", g_ggml_sycl_use_async_mem_op_requested);
382
+
383
+ #ifdef GGML_SYCL_SUPPORT_LEVEL_ZERO_API
384
+ GGML_LOG_INFO(" GGML_SYCL_USE_LEVEL_ZERO_API: %d\n", g_ggml_sycl_use_level_zero_api);
385
+ #else
386
+ GGML_LOG_INFO(" GGML_SYCL_USE_LEVEL_ZERO_API: Disable Level Zero API usage by compile flag\n");
387
+ #endif
388
+
389
+ GGML_LOG_INFO(" GGML_SYCL_USM_SYSTEM: %d\n", g_ggml_sycl_usm_system);
390
+
274
391
  /* NOT REMOVE, keep it for next optimize for XMX.
275
392
  #if defined(SYCL_USE_XMX)
276
393
  fprintf(stderr, "%s: SYCL_USE_XMX: yes\n", __func__);
@@ -278,11 +395,11 @@ static void ggml_check_sycl() try {
278
395
  fprintf(stderr, "%s: SYCL_USE_XMX: no\n", __func__);
279
396
  #endif
280
397
  */
281
- // Currently, we only use async malloc / free when graphs are enabled as it is required for the calls to be
282
- // properly recorded. As this SYCL extension matures it may be beneficial to enable as the default path and in
283
- // other places.
398
+ // Async USM allocation/free is also useful outside the graph path: it avoids the host waits in the reorder
399
+ // staging path while preserving queue ordering semantics. Graph support still depends on the extension being
400
+ // available, but it no longer needs to control the non-graph fast path.
284
401
  #if defined(GGML_SYCL_GRAPH) && SYCL_EXT_ONEAPI_ASYNC_MEMORY_ALLOC
285
- g_ggml_sycl_use_async_mem_op = !g_ggml_sycl_disable_graph;
402
+ g_ggml_sycl_use_async_mem_op = g_ggml_sycl_use_async_mem_op_requested || g_ggml_sycl_enable_graph;
286
403
  if (g_ggml_sycl_use_async_mem_op) {
287
404
  for (unsigned int i = 0; i < dpct::dev_mgr::instance().device_count(); ++i) {
288
405
  if (!dpct::dev_mgr::instance().get_device(i).has(sycl::aspect::ext_oneapi_async_memory_alloc)) {
@@ -346,6 +463,14 @@ catch (sycl::exception const &exc) {
346
463
  std::exit(1);
347
464
  }
348
465
 
466
+ inline void free_aligned_mem_host(void * memblock) {
467
+ #ifdef _WIN32
468
+ _aligned_free(memblock);
469
+ #else
470
+ free(memblock);
471
+ #endif
472
+ }
473
+
349
474
  // sycl buffer
350
475
 
351
476
  struct ggml_backend_sycl_buffer_context {
@@ -355,9 +480,10 @@ struct ggml_backend_sycl_buffer_context {
355
480
  std::string name;
356
481
  optimize_feature opt_feature;
357
482
  std::vector<ggml_tensor_extra_gpu *> tensor_extras;
483
+ bool is_usm_system;
358
484
 
359
- ggml_backend_sycl_buffer_context(int device, void * dev_ptr, queue_ptr stream) :
360
- device(device), dev_ptr(dev_ptr), stream(stream) {
485
+ ggml_backend_sycl_buffer_context(int device, void * dev_ptr, queue_ptr stream, bool is_usm_system) :
486
+ device(device), dev_ptr(dev_ptr), stream(stream), is_usm_system(is_usm_system) {
361
487
  check_allow_gpu_index(device);
362
488
  name = (GGML_SYCL_NAME + std::to_string(device));
363
489
  opt_feature = ggml_sycl_info().devices[device].opt_feature;
@@ -366,7 +492,10 @@ struct ggml_backend_sycl_buffer_context {
366
492
  ~ggml_backend_sycl_buffer_context() {
367
493
  if (dev_ptr != nullptr) {
368
494
  ggml_sycl_set_device(device);
369
- SYCL_CHECK(CHECK_TRY_ERROR(sycl::free(dev_ptr, *stream)));
495
+ if (is_usm_system)
496
+ free_aligned_mem_host(dev_ptr);
497
+ else
498
+ SYCL_CHECK(CHECK_TRY_ERROR(ggml_sycl_free_device(dev_ptr, *stream)));
370
499
  }
371
500
 
372
501
  //release extra used by tensors
@@ -412,11 +541,24 @@ ggml_backend_sycl_buffer_init_tensor(ggml_backend_buffer_t buffer,
412
541
  assert(tensor->view_src->buffer->buft == buffer->buft);
413
542
  return GGML_STATUS_SUCCESS;
414
543
  }
415
- if ((tensor->type == GGML_TYPE_Q4_0 || tensor->type == GGML_TYPE_Q4_K || tensor->type == GGML_TYPE_Q6_K) &&
416
- !g_ggml_sycl_disable_optimize) {
417
- ggml_tensor_extra_gpu * extra = new ggml_tensor_extra_gpu{};
418
- tensor->extra = extra;
419
- ctx->tensor_extras.push_back(extra); //used to release it when destroy ctx.
544
+
545
+ if (g_ggml_sycl_enable_optimize) {
546
+ // set reorder extra buffer based on supported type
547
+ switch (tensor->type) {
548
+ case GGML_TYPE_Q4_0:
549
+ case GGML_TYPE_Q8_0:
550
+ case GGML_TYPE_Q3_K:
551
+ case GGML_TYPE_Q4_K:
552
+ case GGML_TYPE_Q5_K:
553
+ case GGML_TYPE_Q6_K:{
554
+ ggml_tensor_extra_gpu * extra = new ggml_tensor_extra_gpu{};
555
+ tensor->extra = extra;
556
+ ctx->tensor_extras.push_back(extra);
557
+ break;
558
+ }
559
+ default:
560
+ break;
561
+ }
420
562
  }
421
563
 
422
564
  if (ggml_is_quantized(tensor->type)) {
@@ -488,8 +630,50 @@ catch (sycl::exception const &exc) {
488
630
  std::exit(1);
489
631
  }
490
632
 
491
- static void dev2dev_memcpy(sycl::queue &q_dst, sycl::queue &q_src, void *ptr_dst,
633
+ #ifdef GGML_SYCL_SUPPORT_LEVEL_ZERO_API
634
+ static bool ggml_sycl_is_l0_discrete_gpu(int device) {
635
+ return ggml_sycl_info().devices[device].l0_discrete_gpu;
636
+ }
637
+ #endif
638
+
639
+ static void dev2dev_memcpy(int device_dst, sycl::queue &q_dst, int device_src, sycl::queue &q_src, void *ptr_dst,
492
640
  const void *ptr_src, size_t size) {
641
+
642
+ #ifdef GGML_SYCL_SUPPORT_LEVEL_ZERO_API
643
+ if (g_ggml_sycl_dev2dev_memcpy == DEV2DEV_MEMCPY_L0) {
644
+ // Use Level Zero direct copy for dGPU-to-dGPU transfers.
645
+ const bool l0_copy_supported =
646
+ ggml_sycl_is_l0_discrete_gpu(device_dst) && ggml_sycl_is_l0_discrete_gpu(device_src);
647
+ if (g_ggml_sycl_use_level_zero_api && l0_copy_supported) {
648
+ auto ze_ctx = sycl::get_native<sycl::backend::ext_oneapi_level_zero>(q_dst.get_context());
649
+ auto ze_dev = sycl::get_native<sycl::backend::ext_oneapi_level_zero>(q_dst.get_device());
650
+ ze_command_queue_desc_t cq_desc = {ZE_STRUCTURE_TYPE_COMMAND_QUEUE_DESC, nullptr, 0, 0,
651
+ 0, ZE_COMMAND_QUEUE_MODE_SYNCHRONOUS, ZE_COMMAND_QUEUE_PRIORITY_NORMAL};
652
+ ze_command_list_handle_t cl;
653
+ ze_result_t r = zeCommandListCreateImmediate(ze_ctx, ze_dev, &cq_desc, &cl);
654
+ if (r == ZE_RESULT_SUCCESS) {
655
+ GGML_SYCL_DEBUG("[SYCL] dev2dev memcpy by L0\n");
656
+ r = zeCommandListAppendMemoryCopy(cl, ptr_dst, ptr_src, size, nullptr, 0, nullptr);
657
+ zeCommandListDestroy(cl);
658
+ if (r == ZE_RESULT_SUCCESS) {
659
+ return;
660
+ }
661
+ }
662
+ }
663
+ }
664
+ #endif
665
+
666
+ if (g_ggml_sycl_dev2dev_memcpy == DEV2DEV_MEMCPY_SYCL) {
667
+ if (q_dst.get_device().ext_oneapi_can_access_peer(q_src.get_device(),
668
+ sycl::ext::oneapi::peer_access::access_supported)) {
669
+ GGML_SYCL_DEBUG("[SYCL] dev2dev memcpy by SYCL\n");
670
+ SYCL_CHECK(CHECK_TRY_ERROR(q_dst.memcpy(ptr_dst, ptr_src, size).wait()));
671
+ return;
672
+ }
673
+ }
674
+
675
+ // Host-staged copy
676
+ GGML_SYCL_DEBUG("[SYCL] dev2dev memcpy by host forward\n");
493
677
  char *host_buf = (char *)malloc(size);
494
678
  q_src.memcpy(host_buf, (const char *)ptr_src, size).wait();
495
679
  q_dst.memcpy((char *)ptr_dst, host_buf, size).wait();
@@ -536,7 +720,7 @@ ggml_backend_sycl_buffer_cpy_tensor(ggml_backend_buffer_t buffer,
536
720
  size_t size = ggml_nbytes(src);
537
721
 
538
722
  //todo. it's dirty solutino to walkaroud known issue:device2device cross GPUs.
539
- dev2dev_memcpy(*stream_dst, *stream_src, dst->data, src->data, size);
723
+ dev2dev_memcpy(dst_ctx->device, *stream_dst, src_ctx->device, *stream_src, dst->data, src->data, size);
540
724
 
541
725
  //todo, it's known issue:error in device2device cross GPUs. reused when the issue is fixed. DON"T remove
542
726
  #if 0
@@ -570,9 +754,15 @@ static void ggml_backend_sycl_buffer_clear(ggml_backend_buffer_t buffer,
570
754
  SYCL_CHECK(
571
755
  CHECK_TRY_ERROR(dpct::get_current_device().queues_wait_and_throw()));
572
756
 
573
- SYCL_CHECK(CHECK_TRY_ERROR((*stream)
574
- .memset(ctx->dev_ptr, value, buffer->size)
575
- .wait()));
757
+ constexpr size_t MAX_CHUNK = 2ULL << 30; // 2 GiB
758
+ for (size_t off = 0; off < buffer->size; off += MAX_CHUNK) {
759
+ size_t chunk = std::min(buffer->size - off, MAX_CHUNK);
760
+ SYCL_CHECK(CHECK_TRY_ERROR(
761
+ (*stream)
762
+ .memset(static_cast<char*>(ctx->dev_ptr) + off, value, chunk)
763
+ .wait()
764
+ ));
765
+ }
576
766
  }
577
767
  catch (sycl::exception const &exc) {
578
768
  std::cerr << exc.what() << "Exception caught at file:" << __FILE__
@@ -622,6 +812,8 @@ static const ggml_backend_buffer_i ggml_backend_sycl_buffer_interface = {
622
812
  /* .memset_tensor = */ ggml_backend_sycl_buffer_memset_tensor,
623
813
  /* .set_tensor = */ ggml_backend_sycl_buffer_set_tensor,
624
814
  /* .get_tensor = */ ggml_backend_sycl_buffer_get_tensor,
815
+ /* .set_tensor_2d = */ NULL,
816
+ /* .get_tensor_2d = */ NULL,
625
817
  /* .cpy_tensor = */ ggml_backend_sycl_buffer_cpy_tensor,
626
818
  /* .clear = */ ggml_backend_sycl_buffer_clear,
627
819
  /* .reset = */ ggml_backend_sycl_buffer_reset,
@@ -642,22 +834,59 @@ static const char * ggml_backend_sycl_buffer_type_get_name(ggml_backend_buffer_t
642
834
  return ctx->name.c_str();
643
835
  }
644
836
 
837
+ static bool check_usm_system(int device, size_t size) {
838
+ bool use_usm_system = g_ggml_sycl_usm_system && size >= MEM_SIZE_1G;
839
+
840
+ if (use_usm_system && !ggml_sycl_info().devices[device].usm_system_support) {
841
+ GGML_LOG_INFO("Device does not support USM system allocations\n");
842
+ use_usm_system = false;
843
+ }
844
+
845
+ return use_usm_system;
846
+ }
847
+
848
+ inline void * aligned_malloc_host(size_t alignment, size_t size) {
849
+ #ifdef _WIN32
850
+ return _aligned_malloc(size, alignment);
851
+ #else
852
+ return aligned_alloc(alignment, size);
853
+ #endif
854
+ }
855
+
645
856
  static ggml_backend_buffer_t
646
857
  ggml_backend_sycl_buffer_type_alloc_buffer(ggml_backend_buffer_type_t buft,
647
858
  size_t size) try {
859
+ ggml_check_sycl();
860
+
648
861
  ggml_backend_sycl_buffer_type_context * buft_ctx = (ggml_backend_sycl_buffer_type_context *)buft->context;
649
862
  ggml_sycl_set_device(buft_ctx->device);
650
863
  const queue_ptr stream = buft_ctx->stream;
651
864
  size = std::max(size, (size_t)1); // syclMalloc returns null for size 0
865
+ /*
866
+ Alignment below ensures best performance. While in theory it could lead to
867
+ wasting memory, this is acceptable because in practice only few buffers are
868
+ allocated and even less exceed the minimum size accepted here for USM system
869
+ allocations.
870
+ */
871
+ size_t alignment = MEM_SIZE_2M;
872
+ size_t aligned_size = ((size + alignment - 1) / alignment) * alignment;
873
+ bool use_usm_system = check_usm_system(buft_ctx->device, aligned_size);
652
874
 
653
875
  void * dev_ptr;
654
- SYCL_CHECK(CHECK_TRY_ERROR(dev_ptr = (void *)sycl::malloc_device(
655
- size, *stream)));
656
- if (!dev_ptr) {
657
- GGML_LOG_ERROR("%s: can't allocate %lu Bytes of memory on device\n", __func__, size);
658
- return nullptr;
876
+ if (use_usm_system) {
877
+ dev_ptr = (void *)aligned_malloc_host(alignment, aligned_size);
878
+ if (!dev_ptr) {
879
+ GGML_LOG_ERROR("%s: can't allocate %lu Bytes of memory on host\n", __func__, size);
880
+ return nullptr;
881
+ }
882
+ } else {
883
+ SYCL_CHECK(CHECK_TRY_ERROR(dev_ptr = (void *)ggml_sycl_malloc_device(size, *stream)));
884
+ if (!dev_ptr) {
885
+ GGML_LOG_ERROR("%s: can't allocate %lu Bytes of memory on device\n", __func__, size);
886
+ return nullptr;
887
+ }
659
888
  }
660
- ggml_backend_sycl_buffer_context * ctx = new ggml_backend_sycl_buffer_context(buft_ctx->device, dev_ptr, buft_ctx->stream);
889
+ ggml_backend_sycl_buffer_context * ctx = new ggml_backend_sycl_buffer_context(buft_ctx->device, dev_ptr, buft_ctx->stream, use_usm_system);
661
890
  return ggml_backend_buffer_init(buft, ggml_backend_sycl_buffer_interface, ctx, size);
662
891
  }
663
892
  catch (sycl::exception const &exc) {
@@ -667,7 +896,7 @@ catch (sycl::exception const &exc) {
667
896
  }
668
897
 
669
898
  static size_t ggml_backend_sycl_buffer_type_get_alignment(ggml_backend_buffer_type_t buft) {
670
- return 128;
899
+ return SYCL_BUFFER_ALIGNMENT;
671
900
  GGML_UNUSED(buft);
672
901
  }
673
902
 
@@ -775,6 +1004,7 @@ static int64_t get_row_rounding(ggml_type type, const std::array<float, GGML_SYC
775
1004
  }
776
1005
 
777
1006
  switch(type) {
1007
+ case GGML_TYPE_Q1_0:
778
1008
  case GGML_TYPE_Q4_0:
779
1009
  case GGML_TYPE_Q4_1:
780
1010
  return max_compute_capability >= VER_GEN9 ? 128 : 64;
@@ -893,18 +1123,10 @@ ggml_backend_sycl_split_buffer_init_tensor(ggml_backend_buffer_t buffer,
893
1123
  size += ggml_row_size(tensor->type, MATRIX_ROW_PADDING - ne0 % MATRIX_ROW_PADDING);
894
1124
  }
895
1125
 
896
- // FIXME: do not crash if SYCL Buffer alloc fails
897
- // currently, init_tensor cannot fail, it needs to be fixed in ggml-backend first
898
1126
  ggml_sycl_set_device(i);
899
1127
  const queue_ptr stream = ctx->streams[i];
900
1128
  char * buf;
901
- /*
902
- DPCT1009:208: SYCL uses exceptions to report errors and does not use the
903
- error codes. The original code was commented out and a warning string
904
- was inserted. You need to rewrite this code.
905
- */
906
- SYCL_CHECK(CHECK_TRY_ERROR(buf = (char *)sycl::malloc_device(
907
- size, *stream)));
1129
+ SYCL_CHECK(CHECK_TRY_ERROR(buf = (char *)ggml_sycl_malloc_device(size, *stream)));
908
1130
  if (!buf) {
909
1131
  char err_buf[1024];
910
1132
  snprintf(err_buf, 1023, "%s: can't allocate %lu Bytes of memory on device\n", __func__, size);
@@ -1068,6 +1290,8 @@ static struct ggml_backend_buffer_i ggml_backend_sycl_split_buffer_interface = {
1068
1290
  /* .memset_tensor = */ NULL,
1069
1291
  /* .set_tensor = */ ggml_backend_sycl_split_buffer_set_tensor,
1070
1292
  /* .get_tensor = */ ggml_backend_sycl_split_buffer_get_tensor,
1293
+ /* .set_tensor_2d = */ NULL,
1294
+ /* .get_tensor_2d = */ NULL,
1071
1295
  /* .cpy_tensor = */ NULL,
1072
1296
  /* .clear = */ ggml_backend_sycl_split_buffer_clear,
1073
1297
  /* .reset = */ NULL,
@@ -1096,7 +1320,7 @@ static ggml_backend_buffer_t ggml_backend_sycl_split_buffer_type_alloc_buffer(gg
1096
1320
  }
1097
1321
 
1098
1322
  static size_t ggml_backend_sycl_split_buffer_type_get_alignment(ggml_backend_buffer_type_t buft) {
1099
- return 128;
1323
+ return SYCL_BUFFER_ALIGNMENT;
1100
1324
  GGML_UNUSED(buft);
1101
1325
  }
1102
1326
 
@@ -1190,22 +1414,6 @@ static const char * ggml_backend_sycl_host_buffer_type_name(ggml_backend_buffer_
1190
1414
  GGML_UNUSED(buft);
1191
1415
  }
1192
1416
 
1193
- inline void * aligned_malloc_host(size_t alignment, size_t size) {
1194
- #ifdef _WIN32
1195
- return _aligned_malloc(size, alignment);
1196
- #else
1197
- return aligned_alloc(alignment, size);
1198
- #endif
1199
- }
1200
-
1201
- inline void free_aligned_mem_host(void * memblock) {
1202
- #ifdef _WIN32
1203
- _aligned_free(memblock);
1204
- #else
1205
- free(memblock);
1206
- #endif
1207
- }
1208
-
1209
1417
  static void ggml_backend_sycl_host_buffer_free_buffer(ggml_backend_buffer_t buffer) {
1210
1418
  free_aligned_mem_host((void *)buffer->context);
1211
1419
  }
@@ -1260,16 +1468,53 @@ struct ggml_sycl_pool_leg : public ggml_sycl_pool {
1260
1468
  explicit ggml_sycl_pool_leg(queue_ptr qptr_, int device_) : device(device_), qptr(qptr_) {}
1261
1469
 
1262
1470
  ~ggml_sycl_pool_leg() {
1471
+ #ifdef DEBUG_SYCL_POOL
1472
+ int n_cached = 0;
1473
+ size_t bytes_cached = 0;
1474
+ for (int i = 0; i < MAX_SYCL_BUFFERS; ++i) {
1475
+ if (buffer_pool[i].ptr != nullptr) {
1476
+ ++n_cached;
1477
+ bytes_cached += buffer_pool[i].size;
1478
+ }
1479
+ }
1480
+ GGML_LOG_INFO("%s: %d buffers, cached = %.2f MiB\n", __func__,
1481
+ n_cached, bytes_cached / 1024.0 / 1024.0);
1482
+ const auto slots = format_slots_in_alloc_order();
1483
+ if (!slots.empty()) {
1484
+ GGML_LOG_INFO("%s: slots MiB: %s\n", __func__, slots.c_str());
1485
+ }
1486
+ #endif
1487
+
1263
1488
  for (int i = 0; i < MAX_SYCL_BUFFERS; ++i) {
1264
1489
  ggml_sycl_buffer & b = buffer_pool[i];
1265
1490
  if (b.ptr != nullptr) {
1266
- SYCL_CHECK(CHECK_TRY_ERROR(sycl::free(b.ptr, *qptr)));
1491
+ SYCL_CHECK(CHECK_TRY_ERROR(ggml_sycl_free_device(b.ptr, *qptr)));
1267
1492
  pool_size -= b.size;
1268
1493
  }
1269
1494
  }
1270
1495
  GGML_ASSERT(pool_size == 0);
1271
1496
  }
1272
1497
 
1498
+ #ifdef DEBUG_SYCL_POOL
1499
+ std::string format_slots_in_alloc_order() const {
1500
+ std::string line;
1501
+ char buf[32];
1502
+ bool first = true;
1503
+ for (int i = 0; i < MAX_SYCL_BUFFERS; ++i) {
1504
+ if (buffer_pool[i].ptr == nullptr) {
1505
+ continue;
1506
+ }
1507
+ if (!first) {
1508
+ line += '/';
1509
+ }
1510
+ first = false;
1511
+ snprintf(buf, sizeof(buf), "%.2f", buffer_pool[i].size / 1024.0 / 1024.0);
1512
+ line += buf;
1513
+ }
1514
+ return line;
1515
+ }
1516
+ #endif
1517
+
1273
1518
  void * alloc(size_t size, size_t * actual_size) override {
1274
1519
  #ifdef DEBUG_sycl_MALLOC
1275
1520
  int nnz = 0;
@@ -1311,9 +1556,7 @@ struct ggml_sycl_pool_leg : public ggml_sycl_pool {
1311
1556
  void * ptr;
1312
1557
  size_t look_ahead_size = (size_t) (1.05 * size);
1313
1558
 
1314
- SYCL_CHECK(
1315
- CHECK_TRY_ERROR(ptr = (void *)sycl::malloc_device(
1316
- look_ahead_size, *qptr)));
1559
+ SYCL_CHECK(CHECK_TRY_ERROR(ptr = (void *)ggml_sycl_malloc_device(look_ahead_size, *qptr)));
1317
1560
  if (!ptr) {
1318
1561
  GGML_LOG_ERROR("%s: can't allocate %lu Bytes of memory on device/GPU\n", __func__, look_ahead_size);
1319
1562
  return nullptr;
@@ -1341,11 +1584,126 @@ struct ggml_sycl_pool_leg : public ggml_sycl_pool {
1341
1584
  }
1342
1585
  }
1343
1586
  GGML_LOG_WARN("WARNING: sycl buffer pool full, increase MAX_sycl_BUFFERS\n");
1344
- SYCL_CHECK(CHECK_TRY_ERROR(sycl::free(ptr, *qptr)));
1587
+ SYCL_CHECK(CHECK_TRY_ERROR(ggml_sycl_free_device(ptr, *qptr)));
1345
1588
  pool_size -= size;
1346
1589
  }
1347
1590
  };
1348
1591
 
1592
+ // pool with virtual memory management
1593
+ #if defined(GGML_SYCL_SUPPORT_VMM)
1594
+ struct ggml_sycl_pool_vmm : public ggml_sycl_pool {
1595
+ static const size_t SYCL_POOL_VMM_MAX_SIZE = 1ull << 35; // 32 GB
1596
+
1597
+ int device;
1598
+ sycl::context ctx;
1599
+ sycl::device dev;
1600
+
1601
+ uintptr_t pool_addr = 0;
1602
+ size_t pool_used = 0;
1603
+ size_t pool_size = 0;
1604
+ size_t granularity;
1605
+
1606
+ // physical_mem owns the commits (unlike cuMemMap)
1607
+ struct mapping {
1608
+ sycl::ext::oneapi::experimental::physical_mem phys;
1609
+ void * map_ptr;
1610
+ };
1611
+ std::vector<mapping> mappings;
1612
+
1613
+ explicit ggml_sycl_pool_vmm(queue_ptr qptr_, int device_) :
1614
+ device(device_),
1615
+ ctx(qptr_->get_context()),
1616
+ dev(qptr_->get_device()),
1617
+ granularity(ggml_sycl_info().devices[device_].vmm_granularity) {
1618
+ }
1619
+
1620
+ ~ggml_sycl_pool_vmm() {
1621
+ if (pool_addr == 0) {
1622
+ return;
1623
+ }
1624
+
1625
+ // Per spec, unmap must (a) match the exact (ptr, size) of an earlier
1626
+ // physical_mem::map() call and (b) precede destruction of the
1627
+ // physical_mem objects (their dtors won't unmap).
1628
+ for (auto & m : mappings) {
1629
+ SYCL_CHECK(CHECK_TRY_ERROR(sycl::ext::oneapi::experimental::unmap(
1630
+ m.map_ptr, m.phys.size(), ctx)));
1631
+ }
1632
+ SYCL_CHECK(CHECK_TRY_ERROR(sycl::ext::oneapi::experimental::free_virtual_mem(
1633
+ pool_addr, SYCL_POOL_VMM_MAX_SIZE, ctx)));
1634
+ }
1635
+
1636
+ void * alloc(size_t size, size_t * actual_size) override {
1637
+ // round up the allocation size to the alignment to ensure that all allocations are aligned for all data types
1638
+ size = GGML_PAD(size, SYCL_BUFFER_ALIGNMENT);
1639
+
1640
+ size_t avail = pool_size - pool_used;
1641
+
1642
+ if (size > avail) {
1643
+ // round up to the next multiple of the granularity
1644
+ size_t reserve_size = GGML_PAD(size - avail, granularity);
1645
+
1646
+ GGML_ASSERT(pool_size + reserve_size <= SYCL_POOL_VMM_MAX_SIZE);
1647
+
1648
+ // allocate more physical memory
1649
+ std::optional<sycl::ext::oneapi::experimental::physical_mem> phys;
1650
+ SYCL_CHECK(CHECK_TRY_ERROR(phys.emplace(dev, ctx, reserve_size)));
1651
+
1652
+ // reserve virtual address space (if not already reserved)
1653
+ if (pool_addr == 0) {
1654
+ SYCL_CHECK(CHECK_TRY_ERROR(
1655
+ pool_addr = sycl::ext::oneapi::experimental::reserve_virtual_mem(
1656
+ SYCL_POOL_VMM_MAX_SIZE, ctx)));
1657
+ }
1658
+
1659
+ // map at the end of the pool
1660
+ void * map_ptr = nullptr;
1661
+ SYCL_CHECK(CHECK_TRY_ERROR(
1662
+ map_ptr = phys->map(pool_addr + pool_size, reserve_size,
1663
+ sycl::ext::oneapi::experimental::address_access_mode::read_write)));
1664
+
1665
+ // stash these so we could unmap this exact range in dtor
1666
+ mappings.push_back({
1667
+ std::move(*phys),
1668
+ map_ptr,
1669
+ });
1670
+
1671
+ // add to the pool
1672
+ pool_size += reserve_size;
1673
+
1674
+ #ifdef DEBUG_SYCL_MALLOC
1675
+ GGML_LOG_INFO("sycl pool[%d]: size increased to %llu MB (reserved %llu MB)\n",
1676
+ device, (unsigned long long) (pool_size/1024/1024),
1677
+ (unsigned long long) (reserve_size/1024/1024));
1678
+ #endif
1679
+ }
1680
+
1681
+ GGML_ASSERT(pool_addr != 0);
1682
+
1683
+ void * ptr = reinterpret_cast<void *>(pool_addr + pool_used);
1684
+ *actual_size = size;
1685
+ pool_used += size;
1686
+
1687
+ #ifdef DEBUG_SYCL_MALLOC
1688
+ GGML_LOG_INFO("sycl pool[%d]: allocated %llu bytes at %p\n", device, (unsigned long long) size, ptr);
1689
+ #endif
1690
+
1691
+ return ptr;
1692
+ }
1693
+
1694
+ void free(void * ptr, size_t size) override {
1695
+ #ifdef DEBUG_SYCL_MALLOC
1696
+ GGML_LOG_INFO("sycl pool[%d]: freed %llu bytes at %p\n", device, (unsigned long long) size, ptr);
1697
+ #endif
1698
+
1699
+ pool_used -= size;
1700
+
1701
+ // all deallocations must be in reverse order of the allocations
1702
+ GGML_ASSERT(ptr == reinterpret_cast<void *>(pool_addr + pool_used));
1703
+ }
1704
+ };
1705
+ #endif // defined(GGML_SYCL_SUPPORT_VMM)
1706
+
1349
1707
  struct ggml_sycl_pool_host : public ggml_sycl_pool {
1350
1708
  queue_ptr qptr;
1351
1709
  int device;
@@ -1426,15 +1784,18 @@ std::unique_ptr<ggml_sycl_pool> ggml_backend_sycl_context::new_pool_for_host(que
1426
1784
  }
1427
1785
 
1428
1786
  std::unique_ptr<ggml_sycl_pool> ggml_backend_sycl_context::new_pool_for_device(queue_ptr qptr, int device) {
1429
- // TBD: NO VMM support
1430
- // if (ggml_sycl_info().devices[device].vmm) {
1431
- // return std::unique_ptr<ggml_sycl_pool>(new ggml_sycl_pool_vmm(device));
1432
- // }
1433
- return std::unique_ptr<ggml_sycl_pool>(new ggml_sycl_pool_leg(qptr, device));
1787
+ #if defined(GGML_SYCL_SUPPORT_VMM)
1788
+ if (g_ggml_sycl_enable_vmm && ggml_sycl_info().devices[device].vmm) {
1789
+ return std::unique_ptr<ggml_sycl_pool>(new ggml_sycl_pool_vmm(qptr, device));
1790
+ }
1791
+ #endif // defined(GGML_SYCL_SUPPORT_VMM)
1792
+ return std::unique_ptr<ggml_sycl_pool>(new ggml_sycl_pool_leg(qptr, device));
1434
1793
  }
1435
1794
 
1436
- // TBD pool with virtual memory management
1437
- // struct ggml_sycl_pool_vmm : public ggml_sycl_pool
1795
+
1796
+ std::unique_ptr<ggml_sycl_fattn_kv_buffers> ggml_backend_sycl_context::new_fattn_kv_buffers(queue_ptr qptr, int device) {
1797
+ return std::unique_ptr<ggml_sycl_fattn_kv_buffers>(new ggml_sycl_fattn_kv_buffers(qptr, device));
1798
+ }
1438
1799
 
1439
1800
  /// kernels
1440
1801
  typedef void (*ggml_sycl_op_mul_mat_t)(
@@ -1678,69 +2039,6 @@ static void scale_f32(const float * x, float * dst, const float scale, const flo
1678
2039
  }
1679
2040
 
1680
2041
 
1681
- template <typename Ti, typename To>
1682
- static void pool2d_nchw_kernel(
1683
- const int ih, const int iw, const int oh, const int ow,
1684
- const int kh, const int kw, const int sh, const int sw,
1685
- const int ph, const int pw, const int parallel_elements,
1686
- const Ti* src, To* dst, const enum ggml_op_pool op,
1687
- const sycl::nd_item<3> &item_ct1) {
1688
- int idx = item_ct1.get_local_id(2) +
1689
- item_ct1.get_group(2) * item_ct1.get_local_range(2);
1690
- if (idx >= parallel_elements) {
1691
- return;
1692
- }
1693
-
1694
- const int I_HW = ih * iw;
1695
- const int O_HW = oh * ow;
1696
- const int nc = idx / O_HW;
1697
- const int cur_oh = idx % O_HW / ow;
1698
- const int cur_ow = idx % O_HW % ow;
1699
- const Ti* i_ptr = src + nc * I_HW;
1700
- To* o_ptr = dst + nc * O_HW;
1701
- const int start_h = cur_oh * sh - ph;
1702
- const int bh = sycl::max(0, start_h);
1703
- const int eh = sycl::min(ih, start_h + kh);
1704
- const int start_w = cur_ow * sw - pw;
1705
- const int bw = sycl::max(0, start_w);
1706
- const int ew = sycl::min(iw, start_w + kw);
1707
-
1708
- To res = 0;
1709
-
1710
- switch (op) {
1711
- case GGML_OP_POOL_AVG: res = 0; break;
1712
- case GGML_OP_POOL_MAX: res = -FLT_MAX; break;
1713
- default:
1714
- res = (To) sycl::nan(uint32_t(0));
1715
- break;
1716
- }
1717
-
1718
- for (int i = bh; i < eh; i += 1) {
1719
- for (int j = bw; j < ew; j += 1) {
1720
- #if DPCT_COMPATIBILITY_TEMP >= 350
1721
- /*
1722
- DPCT1098:106: The '*' expression is used instead of the __ldg
1723
- call. These two expressions do not provide the exact same
1724
- functionality. Check the generated code for potential precision
1725
- and/or performance issues.
1726
- */
1727
- Ti cur = *(i_ptr + i * iw + j);
1728
- #else
1729
- Ti cur = i_ptr[i * iw + j];
1730
- #endif
1731
- switch (op) {
1732
- case GGML_OP_POOL_AVG: res += (cur / (kh * kw)); break;
1733
- case GGML_OP_POOL_MAX: res = sycl::max(res, (To)cur); break;
1734
- default:
1735
- res = (To) sycl::nan(uint32_t(0));
1736
- break;
1737
- }
1738
- }
1739
- }
1740
- o_ptr[cur_oh * ow + cur_ow] = res;
1741
- }
1742
-
1743
-
1744
2042
  static void ggml_mul_mat_p021_f16_f32_sycl(const void *vx, const float *y,
1745
2043
  float *dst, const int ncols_x,
1746
2044
  const int nrows_x,
@@ -1818,25 +2116,160 @@ static int next_power_of_2(int x) {
1818
2116
  return n;
1819
2117
  }
1820
2118
 
2119
+ static void init_argsort_indices_padded(
2120
+ int * idx,
2121
+ const int nrows,
2122
+ const int ncols_pad,
2123
+ const sycl::nd_item<1> & item_ct1) {
2124
+ const size_t gid = item_ct1.get_local_range(0) * item_ct1.get_group(0) + item_ct1.get_local_id(0);
2125
+ const size_t total = (size_t) nrows * (size_t) ncols_pad;
2126
+
2127
+ if (gid >= total) {
2128
+ return;
2129
+ }
2130
+
2131
+ idx[gid] = (int) (gid % (size_t) ncols_pad);
2132
+ }
2133
+
2134
+ template <ggml_sort_order order>
2135
+ static void argsort_f32_i32_global_pass(const float * x,
2136
+ int * idx,
2137
+ const int ncols,
2138
+ const int nrows,
2139
+ const int ncols_pad,
2140
+ const int j,
2141
+ const int k,
2142
+ const sycl::nd_item<1> & item_ct1) {
2143
+ const size_t gid = item_ct1.get_local_range(0) * item_ct1.get_group(0) + item_ct1.get_local_id(0);
2144
+ const size_t total = (size_t) nrows * (size_t) ncols_pad;
2145
+
2146
+ if (gid >= total) {
2147
+ return;
2148
+ }
2149
+
2150
+ const int row = (int) (gid / (size_t) ncols_pad);
2151
+ const int col = (int) (gid % (size_t) ncols_pad);
2152
+ const int ixj = col ^ j;
2153
+
2154
+ if (ixj <= col || ixj >= ncols_pad) {
2155
+ return;
2156
+ }
2157
+
2158
+ const size_t base = (size_t) row * (size_t) ncols_pad;
2159
+ const size_t pos_a = base + (size_t) col;
2160
+ const size_t pos_b = base + (size_t) ixj;
2161
+
2162
+ const int a = idx[pos_a];
2163
+ const int b = idx[pos_b];
2164
+
2165
+ bool do_swap = false;
2166
+
2167
+ if ((col & k) == 0) {
2168
+ if (a >= ncols ||
2169
+ (b < ncols &&
2170
+ (order == GGML_SORT_ORDER_ASC ?
2171
+ x[(size_t) row * (size_t) ncols + (size_t) a] > x[(size_t) row * (size_t) ncols + (size_t) b] :
2172
+ x[(size_t) row * (size_t) ncols + (size_t) a] < x[(size_t) row * (size_t) ncols + (size_t) b]))) {
2173
+ do_swap = true;
2174
+ }
2175
+ } else {
2176
+ if (b >= ncols ||
2177
+ (a < ncols &&
2178
+ (order == GGML_SORT_ORDER_ASC ?
2179
+ x[(size_t) row * (size_t) ncols + (size_t) a] < x[(size_t) row * (size_t) ncols + (size_t) b] :
2180
+ x[(size_t) row * (size_t) ncols + (size_t) a] > x[(size_t) row * (size_t) ncols + (size_t) b]))) {
2181
+ do_swap = true;
2182
+ }
2183
+ }
2184
+
2185
+ if (do_swap) {
2186
+ idx[pos_a] = b;
2187
+ idx[pos_b] = a;
2188
+ }
2189
+ }
2190
+
2191
+ static void copy_argsort_indices_unpadded(const int * idx_padded,
2192
+ int * dst,
2193
+ const int nrows,
2194
+ const int ncols,
2195
+ const int ncols_pad,
2196
+ const sycl::nd_item<1> & item_ct1) {
2197
+ const size_t gid = item_ct1.get_local_range(0) * item_ct1.get_group(0) + item_ct1.get_local_id(0);
2198
+ const size_t total = (size_t) nrows * (size_t) ncols;
2199
+
2200
+ if (gid >= total) {
2201
+ return;
2202
+ }
2203
+
2204
+ const int row = (int) (gid / (size_t) ncols);
2205
+ const int col = (int) (gid % (size_t) ncols);
2206
+
2207
+ dst[(size_t) row * (size_t) ncols + (size_t) col] = idx_padded[(size_t) row * (size_t) ncols_pad + (size_t) col];
2208
+ }
2209
+
1821
2210
  static void argsort_f32_i32_sycl(const float *x, int *dst, const int ncols,
1822
2211
  const int nrows, ggml_sort_order order,
1823
- queue_ptr stream, int device) {
2212
+ queue_ptr stream, int device, ggml_sycl_pool & pool) {
1824
2213
  // bitonic sort requires ncols to be power of 2
1825
2214
  const int ncols_pad = next_power_of_2(ncols);
2215
+ const size_t shared_mem = (size_t) ncols_pad * sizeof(int);
2216
+ const size_t smpbo = ggml_sycl_info().devices[device].smpbo;
1826
2217
 
1827
- int nth = 1;
1828
- int max_block_size = ggml_sycl_info().max_work_group_sizes[device];
1829
- while (nth < ncols_pad && nth < max_block_size)
1830
- nth *= 2;
1831
- if (nth > max_block_size)
1832
- nth = max_block_size;
2218
+ if (shared_mem > smpbo) {
2219
+ ggml_sycl_pool_alloc<int> idx_padded_alloc(pool, (size_t) nrows * (size_t) ncols_pad);
2220
+ int * idx_padded = idx_padded_alloc.get();
1833
2221
 
1834
- const int tasks_per_thread = ncols_pad / nth;
2222
+ constexpr size_t block_size = 256;
2223
+ const size_t total_padded = (size_t) nrows * (size_t) ncols_pad;
2224
+ const size_t nblocks_padded = (total_padded + block_size - 1) / block_size;
1835
2225
 
1836
- const sycl::range<3> block_dims(1, 1, nth);
1837
- const sycl::range<3> block_nums(1, nrows, 1);
1838
- const size_t shared_mem = ncols_pad * sizeof(int);
1839
- GGML_ASSERT(shared_mem<=ggml_sycl_info().devices[device].smpbo);
2226
+ stream->parallel_for(
2227
+ sycl::nd_range<1>(sycl::range<1>(nblocks_padded * block_size), sycl::range<1>(block_size)),
2228
+ [=](sycl::nd_item<1> item_ct1) { init_argsort_indices_padded(idx_padded, nrows, ncols_pad, item_ct1); });
2229
+
2230
+ for (int k = 2; k <= ncols_pad; k *= 2) {
2231
+ for (int j = k / 2; j > 0; j /= 2) {
2232
+ if (order == GGML_SORT_ORDER_ASC) {
2233
+ stream->parallel_for(
2234
+ sycl::nd_range<1>(sycl::range<1>(nblocks_padded * block_size), sycl::range<1>(block_size)),
2235
+ [=](sycl::nd_item<1> item_ct1) {
2236
+ argsort_f32_i32_global_pass<GGML_SORT_ORDER_ASC>(x, idx_padded, ncols, nrows, ncols_pad, j,
2237
+ k, item_ct1);
2238
+ });
2239
+ } else if (order == GGML_SORT_ORDER_DESC) {
2240
+ stream->parallel_for(
2241
+ sycl::nd_range<1>(sycl::range<1>(nblocks_padded * block_size), sycl::range<1>(block_size)),
2242
+ [=](sycl::nd_item<1> item_ct1) {
2243
+ argsort_f32_i32_global_pass<GGML_SORT_ORDER_DESC>(x, idx_padded, ncols, nrows, ncols_pad, j,
2244
+ k, item_ct1);
2245
+ });
2246
+ } else {
2247
+ GGML_ABORT("invalid sort order");
2248
+ }
2249
+ }
2250
+ }
2251
+
2252
+ const size_t total = (size_t) nrows * (size_t) ncols;
2253
+ const size_t nblocks = (total + block_size - 1) / block_size;
2254
+ stream->parallel_for(sycl::nd_range<1>(sycl::range<1>(nblocks * block_size), sycl::range<1>(block_size)),
2255
+ [=](sycl::nd_item<1> item_ct1) {
2256
+ copy_argsort_indices_unpadded(idx_padded, dst, nrows, ncols, ncols_pad, item_ct1);
2257
+ });
2258
+
2259
+ return;
2260
+ }
2261
+
2262
+ int nth = 1;
2263
+ int max_block_size = ggml_sycl_info().max_work_group_sizes[device];
2264
+ while (nth < ncols_pad && nth < max_block_size)
2265
+ nth *= 2;
2266
+ if (nth > max_block_size)
2267
+ nth = max_block_size;
2268
+
2269
+ const int tasks_per_thread = ncols_pad / nth;
2270
+
2271
+ const sycl::range<3> block_dims(1, 1, nth);
2272
+ const sycl::range<3> block_nums(1, nrows, 1);
1840
2273
 
1841
2274
  if (order == GGML_SORT_ORDER_ASC) {
1842
2275
  stream->submit([&](sycl::handler &cgh) {
@@ -2156,6 +2589,31 @@ inline void ggml_sycl_op_mul_mat_sycl(
2156
2589
  #else
2157
2590
  bool use_fp16 = false;
2158
2591
  #endif
2592
+
2593
+ #if GGML_SYCL_DNNL && defined(GGML_SYCL_HAS_BF16)
2594
+ // Fast path for bf16 src0
2595
+ if (src0->type == GGML_TYPE_BF16 && g_ggml_sycl_enable_dnn && ggml_is_contiguous(src0) &&
2596
+ row_diff == src0->ne[1]) {
2597
+ using bf16_t = sycl::ext::oneapi::bfloat16;
2598
+ ggml_sycl_pool_alloc<bf16_t> src1_as_bf16(ctx.pool(), src1_ncols*ne10);
2599
+ if (src1->type != GGML_TYPE_BF16) {
2600
+ const to_bf16_sycl_t to_bf16_sycl = ggml_get_to_bf16_sycl(src1->type, dst);
2601
+ GGML_ASSERT(to_bf16_sycl != nullptr);
2602
+ to_bf16_sycl(src1_ddf_i, src1_as_bf16.get(), src1_ncols*ne10, stream);
2603
+ } else {
2604
+ stream->memcpy(src1_as_bf16.get(), src1_ddf_i, src1_ncols*ne10*sizeof(bf16_t));
2605
+ }
2606
+ DnnlGemmWrapper::row_gemm(ctx, row_diff, src1_ncols, ne10,
2607
+ src0_dd_i, DnnlGemmWrapper::to_dt<bf16_t>(),
2608
+ src1_as_bf16.get(), DnnlGemmWrapper::to_dt<bf16_t>(),
2609
+ dst_dd_i, DnnlGemmWrapper::to_dt<float>(), stream);
2610
+ GGML_UNUSED(dst);
2611
+ GGML_UNUSED(src1_ddq_i);
2612
+ GGML_UNUSED(src1_padded_row_size);
2613
+ return;
2614
+ }
2615
+ #endif
2616
+
2159
2617
  if ((src0->type == GGML_TYPE_F16 || ggml_is_quantized(src0->type)) && use_fp16 && ggml_is_contiguous(src0) &&
2160
2618
  row_diff == src0->ne[1] && dst->op_params[0] == GGML_PREC_DEFAULT) {
2161
2619
  ggml_sycl_pool_alloc<sycl::half> src0_as_f16(ctx.pool());
@@ -2187,7 +2645,7 @@ inline void ggml_sycl_op_mul_mat_sycl(
2187
2645
  : src1_as_f16.get();
2188
2646
 
2189
2647
  #if GGML_SYCL_DNNL
2190
- if (!g_ggml_sycl_disable_dnn) {
2648
+ if (g_ggml_sycl_enable_dnn) {
2191
2649
  DnnlGemmWrapper::row_gemm(ctx,row_diff, src1_ncols , ne10, src0_ptr,
2192
2650
  DnnlGemmWrapper::to_dt<sycl::half>(), src1_ptr, DnnlGemmWrapper::to_dt<sycl::half>(),
2193
2651
  dst_dd_i, DnnlGemmWrapper::to_dt<float>(), stream);
@@ -2233,21 +2691,25 @@ inline void ggml_sycl_op_mul_mat_sycl(
2233
2691
  const float * src0_ddf_i = src0->type == GGML_TYPE_F32 ? (const float *) src0_dd_i : src0_ddq_as_f32.get();
2234
2692
  const float * src1_ddf1_i = src1->type == GGML_TYPE_F32 ? (const float *) src1_ddf_i : src1_ddq_as_f32.get();
2235
2693
 
2694
+ {
2695
+ const int64_t gemm_flops = (int64_t)row_diff * src1_ncols * ne10;
2696
+ const bool use_mkl_direct = gemm_flops < 256 * 256 * 256;
2236
2697
  #if GGML_SYCL_DNNL
2237
- if (!g_ggml_sycl_disable_dnn) {
2238
- DnnlGemmWrapper::row_gemm(ctx, row_diff, src1_ncols, ne10, src0_ddf_i,
2239
- DnnlGemmWrapper::to_dt<float>(), src1_ddf1_i, DnnlGemmWrapper::to_dt<float>(),
2240
- dst_dd_i, DnnlGemmWrapper::to_dt<float>(), stream);
2241
- }
2242
- else
2698
+ if (g_ggml_sycl_enable_dnn && !use_mkl_direct) {
2699
+ DnnlGemmWrapper::row_gemm(ctx, row_diff, src1_ncols, ne10, src0_ddf_i,
2700
+ DnnlGemmWrapper::to_dt<float>(), src1_ddf1_i, DnnlGemmWrapper::to_dt<float>(),
2701
+ dst_dd_i, DnnlGemmWrapper::to_dt<float>(), stream);
2702
+ }
2703
+ else
2243
2704
  #endif
2244
- {
2245
- const float alpha = 1.0f;
2246
- const float beta = 0.0f;
2247
- SYCL_CHECK(CHECK_TRY_ERROR(oneapi::mkl::blas::column_major::gemm(
2248
- *stream, oneapi::mkl::transpose::trans, oneapi::mkl::transpose::nontrans, row_diff,
2249
- src1_ncols, ne10, dpct::get_value(&alpha, *stream), src0_ddf_i, ne00, src1_ddf1_i, ne10,
2250
- dpct::get_value(&beta, *stream), dst_dd_i, ldc)));
2705
+ {
2706
+ const float alpha = 1.0f;
2707
+ const float beta = 0.0f;
2708
+ SYCL_CHECK(CHECK_TRY_ERROR(oneapi::mkl::blas::column_major::gemm(
2709
+ *stream, oneapi::mkl::transpose::trans, oneapi::mkl::transpose::nontrans, row_diff,
2710
+ src1_ncols, ne10, dpct::get_value(&alpha, *stream), src0_ddf_i, ne00, src1_ddf1_i, ne10,
2711
+ dpct::get_value(&beta, *stream), dst_dd_i, ldc)));
2712
+ }
2251
2713
  }
2252
2714
  }
2253
2715
  GGML_UNUSED(dst);
@@ -2260,45 +2722,6 @@ catch (sycl::exception const &exc) {
2260
2722
  std::exit(1);
2261
2723
  }
2262
2724
 
2263
- static void ggml_sycl_op_pool2d(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
2264
- GGML_ASSERT(dst->src[0]->type == GGML_TYPE_F32);
2265
- GGML_ASSERT( dst->type == GGML_TYPE_F32);
2266
- dpct::queue_ptr main_stream = ctx.stream();
2267
- SYCL_CHECK(ggml_sycl_set_device(ctx.device));
2268
- const float * src0_dd = static_cast<const float *>(dst->src[0]->data);
2269
- float * dst_dd = static_cast<float *>(dst->data);
2270
-
2271
- const int32_t * opts = (const int32_t *)dst->op_params;
2272
- enum ggml_op_pool op = static_cast<ggml_op_pool>(opts[0]);
2273
- const int k0 = opts[1];
2274
- const int k1 = opts[2];
2275
- const int s0 = opts[3];
2276
- const int s1 = opts[4];
2277
- const int p0 = opts[5];
2278
- const int p1 = opts[6];
2279
-
2280
- const int64_t IH = dst->src[0]->ne[1];
2281
- const int64_t IW = dst->src[0]->ne[0];
2282
-
2283
- const int64_t N = dst->ne[3];
2284
- const int64_t OC = dst->ne[2];
2285
- const int64_t OH = dst->ne[1];
2286
- const int64_t OW = dst->ne[0];
2287
-
2288
- const int parallel_elements = N * OC * OH * OW;
2289
- const int num_blocks = (parallel_elements + SYCL_POOL2D_BLOCK_SIZE - 1) / SYCL_POOL2D_BLOCK_SIZE;
2290
- sycl::range<3> block_nums(1, 1, num_blocks);
2291
- main_stream->parallel_for(
2292
- sycl::nd_range<3>(block_nums *
2293
- sycl::range<3>(1, 1, SYCL_IM2COL_BLOCK_SIZE),
2294
- sycl::range<3>(1, 1, SYCL_IM2COL_BLOCK_SIZE)),
2295
- [=](sycl::nd_item<3> item_ct1) {
2296
- pool2d_nchw_kernel(IH, IW, OH, OW, k1, k0, s1, s0, p1, p0,
2297
- parallel_elements, src0_dd, dst_dd, op,
2298
- item_ct1);
2299
- });
2300
- }
2301
-
2302
2725
  inline void ggml_sycl_op_sum(ggml_backend_sycl_context & ctx, ggml_tensor *dst) {
2303
2726
  GGML_ASSERT(dst->src[0]->type == GGML_TYPE_F32);
2304
2727
  GGML_ASSERT( dst->type == GGML_TYPE_F32);
@@ -2365,7 +2788,7 @@ inline void ggml_sycl_op_argsort(ggml_backend_sycl_context & ctx, ggml_tensor *
2365
2788
  enum ggml_sort_order order = (enum ggml_sort_order) dst->op_params[0];
2366
2789
 
2367
2790
  argsort_f32_i32_sycl(src0_dd, (int *)dst_dd, ncols, nrows, order,
2368
- main_stream, ctx.device);
2791
+ main_stream, ctx.device, ctx.pool());
2369
2792
  }
2370
2793
 
2371
2794
  static void ggml_sycl_op_top_k(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
@@ -2758,7 +3181,7 @@ static void ggml_sycl_op_mul_mat(ggml_backend_sycl_context & ctx, const ggml_ten
2758
3181
  src1_ddf_i_source += (i0 * ne11 + src1_col_0) * ne10;
2759
3182
 
2760
3183
  SYCL_CHECK(
2761
- CHECK_TRY_ERROR(dev2dev_memcpy(*stream, *main_stream, src1_ddf_i, src1_ddf_i_source,
3184
+ CHECK_TRY_ERROR(dev2dev_memcpy(i, *stream, ctx.device, *main_stream, src1_ddf_i, src1_ddf_i_source,
2762
3185
  src1_ncols * ne10 * sizeof(float))));
2763
3186
  }
2764
3187
  }
@@ -3092,7 +3515,7 @@ static void ggml_sycl_mul_mat_batched_sycl(ggml_backend_sycl_context & ctx, cons
3092
3515
  const int64_t r3 = ne13 / ne03;
3093
3516
 
3094
3517
  #if GGML_SYCL_DNNL
3095
- if (!g_ggml_sycl_disable_dnn) {
3518
+ if (g_ggml_sycl_enable_dnn) {
3096
3519
  int64_t str_a0 = nb00 / type_size_src0;
3097
3520
  int64_t str_a1 = nb01 / type_size_src0;
3098
3521
  int64_t str_a2 = nb02 / type_size_src0;
@@ -3248,9 +3671,13 @@ inline bool ggml_sycl_supports_mmq(enum ggml_type type) {
3248
3671
 
3249
3672
  inline bool ggml_sycl_supports_reorder_mul_mat_sycl(enum ggml_type type) {
3250
3673
  switch (type) {
3674
+ case GGML_TYPE_Q1_0:
3251
3675
  case GGML_TYPE_Q4_0:
3676
+ case GGML_TYPE_Q8_0:
3252
3677
  return true;
3678
+ case GGML_TYPE_Q3_K:
3253
3679
  case GGML_TYPE_Q4_K:
3680
+ case GGML_TYPE_Q5_K:
3254
3681
  case GGML_TYPE_Q6_K:
3255
3682
  return !g_ggml_sycl_prioritize_dmmv;
3256
3683
  default:
@@ -3260,7 +3687,13 @@ inline bool ggml_sycl_supports_reorder_mul_mat_sycl(enum ggml_type type) {
3260
3687
 
3261
3688
  inline bool ggml_sycl_supports_reorder_dmmv(enum ggml_type type) {
3262
3689
  switch (type) {
3690
+ case GGML_TYPE_Q1_0:
3263
3691
  case GGML_TYPE_Q4_0:
3692
+ case GGML_TYPE_Q8_0:
3693
+ case GGML_TYPE_Q3_K:
3694
+ case GGML_TYPE_Q4_K:
3695
+ case GGML_TYPE_Q5_K:
3696
+ case GGML_TYPE_Q6_K:
3264
3697
  return true;
3265
3698
  default:
3266
3699
  return false;
@@ -3269,8 +3702,12 @@ inline bool ggml_sycl_supports_reorder_dmmv(enum ggml_type type) {
3269
3702
 
3270
3703
  inline bool ggml_sycl_supports_reorder_mmvq(enum ggml_type type) {
3271
3704
  switch (type) {
3705
+ case GGML_TYPE_Q1_0:
3272
3706
  case GGML_TYPE_Q4_0:
3707
+ case GGML_TYPE_Q8_0:
3708
+ case GGML_TYPE_Q3_K:
3273
3709
  case GGML_TYPE_Q4_K:
3710
+ case GGML_TYPE_Q5_K:
3274
3711
  case GGML_TYPE_Q6_K:
3275
3712
  return true;
3276
3713
  default:
@@ -3280,6 +3717,7 @@ inline bool ggml_sycl_supports_reorder_mmvq(enum ggml_type type) {
3280
3717
 
3281
3718
  static bool ggml_sycl_supports_dmmv(enum ggml_type type) {
3282
3719
  switch (type) {
3720
+ case GGML_TYPE_Q1_0:
3283
3721
  case GGML_TYPE_Q4_0:
3284
3722
  case GGML_TYPE_Q4_1:
3285
3723
  case GGML_TYPE_Q5_0:
@@ -3291,6 +3729,7 @@ static bool ggml_sycl_supports_dmmv(enum ggml_type type) {
3291
3729
  case GGML_TYPE_Q5_K:
3292
3730
  case GGML_TYPE_Q6_K:
3293
3731
  case GGML_TYPE_F16:
3732
+ case GGML_TYPE_BF16:
3294
3733
  return true;
3295
3734
  default:
3296
3735
  return false;
@@ -3308,7 +3747,7 @@ static inline void * sycl_ext_malloc_device(dpct::queue_ptr stream, size_t size)
3308
3747
  // If async allocation extension is not available, use_async should always be false.
3309
3748
  GGML_ASSERT(!use_async);
3310
3749
  #endif
3311
- return sycl::malloc(size, *stream, sycl::usm::alloc::device);
3750
+ return ggml_sycl_malloc_device(size, *stream);
3312
3751
  }
3313
3752
 
3314
3753
  static inline void sycl_ext_free(dpct::queue_ptr stream, void * ptr) {
@@ -3322,12 +3761,58 @@ static inline void sycl_ext_free(dpct::queue_ptr stream, void * ptr) {
3322
3761
  // If async allocation extension is not available, use_async should always be false.
3323
3762
  GGML_ASSERT(!use_async);
3324
3763
  #endif
3325
- sycl::free(ptr, *stream);
3764
+ ggml_sycl_free_device(ptr, *stream);
3326
3765
  }
3327
3766
 
3328
- static void reorder_qw_q4_0(uint8_t * data_device, const int ncols, const int nrows, size_t size, size_t offset,
3767
+ // RAII wrapper for temporary reorder buffers with optional host memory fallback.
3768
+ // When device allocation fails and GGML_SYCL_HOST_MEM_FALLBACK is enabled,
3769
+ // falls back to host memory so the reorder kernel can still run (over PCIe).
3770
+ // Device access to host memory requires Linux kernel 6.8+ (Ubuntu 26.04+).
3771
+ struct sycl_reorder_temp_buffer {
3772
+ void * ptr = nullptr;
3773
+ dpct::queue_ptr stream;
3774
+
3775
+ sycl_reorder_temp_buffer(dpct::queue_ptr stream, size_t size) : stream(stream) {
3776
+ ptr = sycl_ext_malloc_device(stream, size);
3777
+ #ifdef GGML_SYCL_HOST_MEM_FALLBACK
3778
+ if (!ptr) {
3779
+ ptr = sycl::malloc_host(size, *stream);
3780
+ if (ptr) {
3781
+ host_fallback = true;
3782
+ GGML_LOG_WARN("%s: device alloc of %zu bytes failed, using host memory fallback\n", __func__, size);
3783
+ }
3784
+ }
3785
+ #endif
3786
+ }
3787
+
3788
+ ~sycl_reorder_temp_buffer() {
3789
+ if (!ptr) {
3790
+ return;
3791
+ }
3792
+ if (host_fallback) {
3793
+ sycl::free(ptr, *stream);
3794
+ } else {
3795
+ sycl_ext_free(stream, ptr);
3796
+ }
3797
+ }
3798
+
3799
+ explicit operator bool() const { return ptr != nullptr; }
3800
+
3801
+ sycl_reorder_temp_buffer(const sycl_reorder_temp_buffer &) = delete;
3802
+ sycl_reorder_temp_buffer & operator=(const sycl_reorder_temp_buffer &) = delete;
3803
+
3804
+ private:
3805
+ bool host_fallback = false;
3806
+ };
3807
+
3808
+ static bool reorder_qw_q4_0(uint8_t * data_device, const int ncols, const int nrows, size_t size, size_t offset,
3329
3809
  dpct::queue_ptr stream) {
3330
- uint8_t * tmp_buf = static_cast<uint8_t *>(sycl_ext_malloc_device(stream, size));
3810
+ sycl_reorder_temp_buffer tmp(stream, size);
3811
+ if (!tmp) {
3812
+ GGML_LOG_WARN("%s: failed to allocate %zu bytes for reorder temp buffer, skipping reorder\n", __func__, size);
3813
+ return false;
3814
+ }
3815
+ uint8_t * tmp_buf = static_cast<uint8_t *>(tmp.ptr);
3331
3816
 
3332
3817
  sycl::event copy_event;
3333
3818
  SYCL_CHECK(CHECK_TRY_ERROR(copy_event = stream->memcpy(tmp_buf, data_device, size)));
@@ -3356,16 +3841,60 @@ static void reorder_qw_q4_0(uint8_t * data_device, const int ncols, const int nr
3356
3841
  if (!g_ggml_sycl_use_async_mem_op) {
3357
3842
  reorder_event.wait_and_throw();
3358
3843
  }
3359
- sycl_ext_free(stream, tmp_buf);
3844
+ return true;
3845
+ }
3846
+
3847
+ static bool reorder_qw_q8_0(uint8_t * data_device, const int ncols, const int nrows, size_t size, size_t offset,
3848
+ dpct::queue_ptr stream) {
3849
+ sycl_reorder_temp_buffer tmp(stream, size);
3850
+ if (!tmp) {
3851
+ GGML_LOG_WARN("%s: failed to allocate %zu bytes for reorder temp buffer, skipping reorder\n", __func__, size);
3852
+ return false;
3853
+ }
3854
+ uint8_t * tmp_buf = static_cast<uint8_t *>(tmp.ptr);
3855
+
3856
+ sycl::event copy_event;
3857
+ SYCL_CHECK(CHECK_TRY_ERROR(copy_event = stream->memcpy(tmp_buf, data_device, size)));
3858
+ if (!g_ggml_sycl_use_async_mem_op) {
3859
+ copy_event.wait();
3860
+ }
3861
+
3862
+ GGML_ASSERT((size % sizeof(block_q8_0) == 0));
3863
+ GGML_ASSERT((offset % sizeof(block_q8_0) == 0));
3864
+ int offset_blks = offset / sizeof(block_q8_0);
3865
+ auto qs_ptr = data_device + offset_blks * QK8_0;
3866
+ auto d_ptr = (sycl::half*)(qs_ptr + ncols * nrows) + offset_blks;
3867
+
3868
+ auto reorder_event = stream->parallel_for(
3869
+ size / sizeof(block_q8_0),
3870
+ [=](auto i) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
3871
+ const block_q8_0* x = (const block_q8_0*)tmp_buf;
3872
+ const int ib = i;
3873
+
3874
+ for (int j = 0; j < QK8_0; j++)
3875
+ {
3876
+ *((int8_t*)qs_ptr + ib * QK8_0 + j) = x[ib].qs[j];
3877
+ }
3878
+ *(d_ptr + ib) = x[ib].d;
3879
+ });
3880
+ if (!g_ggml_sycl_use_async_mem_op) {
3881
+ reorder_event.wait_and_throw();
3882
+ }
3883
+ return true;
3360
3884
  }
3361
3885
 
3362
- static void reorder_qw_q4_k(uint8_t * data_device, size_t size, size_t offset, dpct::queue_ptr stream) {
3886
+ static bool reorder_qw_q4_k(uint8_t * data_device, size_t size, size_t offset, dpct::queue_ptr stream) {
3363
3887
  GGML_ASSERT(size % sizeof(block_q4_K) == 0);
3364
3888
  GGML_ASSERT(offset % sizeof(block_q4_K) == 0);
3365
3889
 
3366
3890
  const int nblocks = size / sizeof(block_q4_K);
3367
3891
 
3368
- uint8_t * tmp_buf = static_cast<uint8_t *>(sycl_ext_malloc_device(stream, size));
3892
+ sycl_reorder_temp_buffer tmp(stream, size);
3893
+ if (!tmp) {
3894
+ GGML_LOG_WARN("%s: failed to allocate %zu bytes for reorder temp buffer, skipping reorder\n", __func__, size);
3895
+ return false;
3896
+ }
3897
+ uint8_t * tmp_buf = static_cast<uint8_t *>(tmp.ptr);
3369
3898
 
3370
3899
  sycl::event copy_event;
3371
3900
  SYCL_CHECK(CHECK_TRY_ERROR(copy_event = stream->memcpy(tmp_buf, data_device, size)));
@@ -3394,16 +3923,260 @@ static void reorder_qw_q4_k(uint8_t * data_device, size_t size, size_t offset, d
3394
3923
  if (!g_ggml_sycl_use_async_mem_op) {
3395
3924
  reorder_event.wait_and_throw();
3396
3925
  }
3397
- sycl_ext_free(stream, tmp_buf);
3926
+ return true;
3927
+ }
3928
+
3929
+ // Reorder each expert slice into a self-contained SoA layout.
3930
+ static bool reorder_qw_q4_k_moe(uint8_t * data_device, size_t expert_bytes, int64_t n_expert, dpct::queue_ptr stream) {
3931
+ GGML_ASSERT(expert_bytes % sizeof(block_q4_K) == 0);
3932
+ const int blocks_per_expert = (int) (expert_bytes / sizeof(block_q4_K));
3933
+ const size_t total_bytes = expert_bytes * (size_t) n_expert;
3934
+
3935
+ sycl_reorder_temp_buffer tmp(stream, total_bytes);
3936
+ if (!tmp) {
3937
+ GGML_LOG_WARN("%s: failed to allocate %zu bytes for reorder temp buffer, skipping reorder\n", __func__, total_bytes);
3938
+ return false;
3939
+ }
3940
+ uint8_t * tmp_buf = static_cast<uint8_t *>(tmp.ptr);
3941
+
3942
+ sycl::event copy_event;
3943
+ SYCL_CHECK(CHECK_TRY_ERROR(copy_event = stream->memcpy(tmp_buf, data_device, total_bytes)));
3944
+ if (!g_ggml_sycl_use_async_mem_op) {
3945
+ copy_event.wait();
3946
+ }
3947
+
3948
+ const int total_blocks = blocks_per_expert * (int) n_expert;
3949
+ auto reorder_event = stream->parallel_for(total_blocks, [=](auto gb_) {
3950
+ const int gb = gb_;
3951
+ const int e = gb / blocks_per_expert;
3952
+ const int ib = gb % blocks_per_expert;
3953
+ const block_q4_K * x = (const block_q4_K *) (tmp_buf + (size_t) e * expert_bytes);
3954
+ uint8_t * base = data_device + (size_t) e * expert_bytes;
3955
+
3956
+ auto * qs_ptr = base;
3957
+ auto * scales_ptr = qs_ptr + QK_K / 2 * blocks_per_expert;
3958
+ auto * dm_ptr = (sycl::half2 *) (scales_ptr + K_SCALE_SIZE * blocks_per_expert);
3959
+
3960
+ for (int j = 0; j < QK_K / 2; ++j) {
3961
+ qs_ptr[ib * (QK_K / 2) + j] = x[ib].qs[j];
3962
+ }
3963
+ for (int j = 0; j < K_SCALE_SIZE; ++j) {
3964
+ scales_ptr[ib * K_SCALE_SIZE + j] = x[ib].scales[j];
3965
+ }
3966
+ dm_ptr[ib] = x[ib].dm;
3967
+ });
3968
+ if (!g_ggml_sycl_use_async_mem_op) {
3969
+ reorder_event.wait_and_throw();
3970
+ }
3971
+ return true;
3972
+ }
3973
+
3974
+ // Reorder each Q5_K expert slice into [qs][qh][scales][dm].
3975
+ static bool reorder_qw_q5_k_moe(uint8_t * data_device, size_t expert_bytes, int64_t n_expert, dpct::queue_ptr stream) {
3976
+ GGML_ASSERT(expert_bytes % sizeof(block_q5_K) == 0);
3977
+ const int blocks_per_expert = (int) (expert_bytes / sizeof(block_q5_K));
3978
+ const size_t total_bytes = expert_bytes * (size_t) n_expert;
3979
+
3980
+ sycl_reorder_temp_buffer tmp(stream, total_bytes);
3981
+ if (!tmp) {
3982
+ GGML_LOG_WARN("%s: failed to allocate %zu bytes for reorder temp buffer, skipping reorder\n", __func__, total_bytes);
3983
+ return false;
3984
+ }
3985
+ uint8_t * tmp_buf = static_cast<uint8_t *>(tmp.ptr);
3986
+
3987
+ sycl::event copy_event;
3988
+ SYCL_CHECK(CHECK_TRY_ERROR(copy_event = stream->memcpy(tmp_buf, data_device, total_bytes)));
3989
+ if (!g_ggml_sycl_use_async_mem_op) {
3990
+ copy_event.wait();
3991
+ }
3992
+
3993
+ const int total_blocks = blocks_per_expert * (int) n_expert;
3994
+ auto reorder_event = stream->parallel_for(total_blocks, [=](auto gb_) {
3995
+ const int gb = gb_;
3996
+ const int e = gb / blocks_per_expert;
3997
+ const int ib = gb % blocks_per_expert;
3998
+ const block_q5_K * x = (const block_q5_K *) (tmp_buf + (size_t) e * expert_bytes);
3999
+ uint8_t * base = data_device + (size_t) e * expert_bytes;
4000
+
4001
+ auto * qs_ptr = base;
4002
+ auto * qh_ptr = qs_ptr + (QK_K / 2) * blocks_per_expert;
4003
+ auto * scales_ptr = qh_ptr + (QK_K / 8) * blocks_per_expert;
4004
+ auto * dm_ptr = (sycl::half2 *) (scales_ptr + K_SCALE_SIZE * blocks_per_expert);
4005
+
4006
+ for (int j = 0; j < QK_K / 2; ++j) {
4007
+ qs_ptr[ib * (QK_K / 2) + j] = x[ib].qs[j];
4008
+ }
4009
+ for (int j = 0; j < QK_K / 8; ++j) {
4010
+ qh_ptr[ib * (QK_K / 8) + j] = x[ib].qh[j];
4011
+ }
4012
+ for (int j = 0; j < K_SCALE_SIZE; ++j) {
4013
+ scales_ptr[ib * K_SCALE_SIZE + j] = x[ib].scales[j];
4014
+ }
4015
+ dm_ptr[ib] = x[ib].dm;
4016
+ });
4017
+ if (!g_ggml_sycl_use_async_mem_op) {
4018
+ reorder_event.wait_and_throw();
4019
+ }
4020
+ return true;
4021
+ }
4022
+
4023
+ // Reorder each Q6_K expert slice into [ql][qh][scales][d].
4024
+ static bool reorder_qw_q6_k_moe(uint8_t * data_device, size_t expert_bytes, int64_t n_expert, dpct::queue_ptr stream) {
4025
+ GGML_ASSERT(expert_bytes % sizeof(block_q6_K) == 0);
4026
+ const int blocks_per_expert = (int) (expert_bytes / sizeof(block_q6_K));
4027
+ const size_t total_bytes = expert_bytes * (size_t) n_expert;
4028
+
4029
+ sycl_reorder_temp_buffer tmp(stream, total_bytes);
4030
+ if (!tmp) {
4031
+ GGML_LOG_WARN("%s: failed to allocate %zu bytes for reorder temp buffer, skipping reorder\n", __func__, total_bytes);
4032
+ return false;
4033
+ }
4034
+ uint8_t * tmp_buf = static_cast<uint8_t *>(tmp.ptr);
4035
+
4036
+ sycl::event copy_event;
4037
+ SYCL_CHECK(CHECK_TRY_ERROR(copy_event = stream->memcpy(tmp_buf, data_device, total_bytes)));
4038
+ if (!g_ggml_sycl_use_async_mem_op) {
4039
+ copy_event.wait();
4040
+ }
4041
+
4042
+ const int total_blocks = blocks_per_expert * (int) n_expert;
4043
+ auto reorder_event = stream->parallel_for(total_blocks, [=](auto gb_) {
4044
+ const int gb = gb_;
4045
+ const int e = gb / blocks_per_expert;
4046
+ const int ib = gb % blocks_per_expert;
4047
+ const block_q6_K * x = (const block_q6_K *) (tmp_buf + (size_t) e * expert_bytes);
4048
+ uint8_t * base = data_device + (size_t) e * expert_bytes;
4049
+
4050
+ auto * ql_ptr = base;
4051
+ auto * qh_ptr = ql_ptr + (QK_K / 2) * blocks_per_expert;
4052
+ auto * scales_ptr = qh_ptr + (QK_K / 4) * blocks_per_expert;
4053
+ auto * d_ptr = (sycl::half *) (scales_ptr + (QK_K / 16) * blocks_per_expert);
4054
+
4055
+ for (int j = 0; j < QK_K / 2; ++j) {
4056
+ ql_ptr[ib * (QK_K / 2) + j] = x[ib].ql[j];
4057
+ }
4058
+ for (int j = 0; j < QK_K / 4; ++j) {
4059
+ qh_ptr[ib * (QK_K / 4) + j] = x[ib].qh[j];
4060
+ }
4061
+ for (int j = 0; j < QK_K / 16; ++j) {
4062
+ scales_ptr[ib * (QK_K / 16) + j] = x[ib].scales[j];
4063
+ }
4064
+ d_ptr[ib] = x[ib].d;
4065
+ });
4066
+ if (!g_ggml_sycl_use_async_mem_op) {
4067
+ reorder_event.wait_and_throw();
4068
+ }
4069
+ return true;
4070
+ }
4071
+
4072
+ static bool reorder_qw_q3_k(uint8_t * data_device, size_t size, size_t offset, dpct::queue_ptr stream) {
4073
+ GGML_ASSERT(size % sizeof(block_q3_K) == 0);
4074
+ GGML_ASSERT(offset % sizeof(block_q3_K) == 0);
4075
+
4076
+ const int nblocks = size / sizeof(block_q3_K);
4077
+
4078
+ sycl_reorder_temp_buffer tmp(stream, size);
4079
+ if (!tmp) {
4080
+ GGML_LOG_WARN("%s: failed to allocate %zu bytes for reorder temp buffer, skipping reorder\n", __func__, size);
4081
+ return false;
4082
+ }
4083
+ uint8_t * tmp_buf = static_cast<uint8_t *>(tmp.ptr);
4084
+
4085
+ sycl::event copy_event;
4086
+ SYCL_CHECK(CHECK_TRY_ERROR(copy_event = stream->memcpy(tmp_buf, data_device, size)));
4087
+ if (!g_ggml_sycl_use_async_mem_op) {
4088
+ copy_event.wait();
4089
+ }
4090
+
4091
+ auto * qs_ptr = data_device;
4092
+ auto * hmask_ptr = qs_ptr + (QK_K / 4) * nblocks;
4093
+ auto * scales_ptr = hmask_ptr + (QK_K / 8) * nblocks;
4094
+ sycl::half * d_ptr = (sycl::half *) (scales_ptr + 12 * nblocks);
4095
+
4096
+ auto reorder_event = stream->parallel_for(nblocks, [=](auto i) {
4097
+ const block_q3_K * x = (const block_q3_K *) tmp_buf;
4098
+ const int ib = i;
4099
+
4100
+ for (int j = 0; j < QK_K / 4; ++j) {
4101
+ qs_ptr[ib * (QK_K / 4) + j] = x[ib].qs[j];
4102
+ }
4103
+
4104
+ for (int j = 0; j < QK_K / 8; ++j) {
4105
+ hmask_ptr[ib * (QK_K / 8) + j] = x[ib].hmask[j];
4106
+ }
4107
+
4108
+ for (int j = 0; j < 12; ++j) {
4109
+ scales_ptr[ib * 12 + j] = x[ib].scales[j];
4110
+ }
4111
+
4112
+ d_ptr[ib] = x[ib].d;
4113
+ });
4114
+ if (!g_ggml_sycl_use_async_mem_op) {
4115
+ reorder_event.wait_and_throw();
4116
+ }
4117
+ return true;
4118
+ }
4119
+
4120
+ static bool reorder_qw_q5_k(uint8_t * data_device, size_t size, size_t offset, dpct::queue_ptr stream) {
4121
+ GGML_ASSERT(size % sizeof(block_q5_K) == 0);
4122
+ GGML_ASSERT(offset % sizeof(block_q5_K) == 0);
4123
+
4124
+ const int nblocks = size / sizeof(block_q5_K);
4125
+
4126
+ sycl_reorder_temp_buffer tmp(stream, size);
4127
+ if (!tmp) {
4128
+ GGML_LOG_WARN("%s: failed to allocate %zu bytes for reorder temp buffer, skipping reorder\n", __func__, size);
4129
+ return false;
4130
+ }
4131
+ uint8_t * tmp_buf = static_cast<uint8_t *>(tmp.ptr);
4132
+
4133
+ sycl::event copy_event;
4134
+ SYCL_CHECK(CHECK_TRY_ERROR(copy_event = stream->memcpy(tmp_buf, data_device, size)));
4135
+ if (!g_ggml_sycl_use_async_mem_op) {
4136
+ copy_event.wait();
4137
+ }
4138
+
4139
+ auto * qs_ptr = data_device;
4140
+ auto * qh_ptr = qs_ptr + (QK_K / 2) * nblocks;
4141
+ auto * scales_ptr = qh_ptr + (QK_K / 8) * nblocks;
4142
+ auto * dm_ptr = (sycl::half2 *) (scales_ptr + K_SCALE_SIZE * nblocks);
4143
+
4144
+ auto reorder_event = stream->parallel_for(nblocks, [=](auto i) {
4145
+ const block_q5_K * x = (const block_q5_K *) tmp_buf;
4146
+ const int ib = i;
4147
+
4148
+ for (int j = 0; j < QK_K / 2; ++j) {
4149
+ qs_ptr[ib * (QK_K / 2) + j] = x[ib].qs[j];
4150
+ }
4151
+
4152
+ for (int j = 0; j < QK_K / 8; ++j) {
4153
+ qh_ptr[ib * (QK_K / 8) + j] = x[ib].qh[j];
4154
+ }
4155
+
4156
+ for (int j = 0; j < K_SCALE_SIZE; ++j) {
4157
+ scales_ptr[ib * K_SCALE_SIZE + j] = x[ib].scales[j];
4158
+ }
4159
+
4160
+ dm_ptr[ib] = x[ib].dm;
4161
+ });
4162
+ if (!g_ggml_sycl_use_async_mem_op) {
4163
+ reorder_event.wait_and_throw();
4164
+ }
4165
+ return true;
3398
4166
  }
3399
4167
 
3400
- static void reorder_qw_q6_k(uint8_t * data_device, size_t size, size_t offset, dpct::queue_ptr stream) {
4168
+ static bool reorder_qw_q6_k(uint8_t * data_device, size_t size, size_t offset, dpct::queue_ptr stream) {
3401
4169
  GGML_ASSERT(size % sizeof(block_q6_K) == 0);
3402
4170
  GGML_ASSERT(offset % sizeof(block_q6_K) == 0);
3403
4171
 
3404
4172
  const int nblocks = size / sizeof(block_q6_K);
3405
4173
 
3406
- uint8_t * tmp_buf = static_cast<uint8_t *>(sycl_ext_malloc_device(stream, size));
4174
+ sycl_reorder_temp_buffer tmp(stream, size);
4175
+ if (!tmp) {
4176
+ GGML_LOG_WARN("%s: failed to allocate %zu bytes for reorder temp buffer, skipping reorder\n", __func__, size);
4177
+ return false;
4178
+ }
4179
+ uint8_t * tmp_buf = static_cast<uint8_t *>(tmp.ptr);
3407
4180
 
3408
4181
  sycl::event copy_event;
3409
4182
  SYCL_CHECK(CHECK_TRY_ERROR(copy_event = stream->memcpy(tmp_buf, data_device, size)));
@@ -3442,36 +4215,56 @@ static void reorder_qw_q6_k(uint8_t * data_device, size_t size, size_t offset, d
3442
4215
  if (!g_ggml_sycl_use_async_mem_op) {
3443
4216
  reorder_event.wait_and_throw();
3444
4217
  }
3445
- sycl_ext_free(stream, tmp_buf);
4218
+ return true;
3446
4219
  }
3447
4220
 
3448
- static void reorder_qw(const ggml_tensor * src0, dpct::queue_ptr stream) {
4221
+ static bool reorder_qw(const ggml_tensor * src0, dpct::queue_ptr stream) {
3449
4222
  uint8_t * data_device = (uint8_t *) src0->data;
3450
4223
  size_t ncols = src0->ne[0];
3451
4224
  size_t nrows = src0->ne[1];
3452
4225
  size_t size = ggml_nbytes(src0);
3453
4226
 
3454
- switch (src0->type) {
3455
- case GGML_TYPE_Q4_0:
3456
- reorder_qw_q4_0(data_device, ncols, nrows, size, 0, stream);
3457
- break;
4227
+ // MoE expert weights are addressed per expert via nb[2], so each slice must
4228
+ // remain self-contained after reorder.
4229
+ if (src0->ne[2] > 1) {
4230
+ GGML_ASSERT((size_t) size == (size_t) src0->ne[2] * src0->nb[2]);
4231
+ switch (src0->type) {
4232
+ case GGML_TYPE_Q4_K:
4233
+ return reorder_qw_q4_k_moe(data_device, src0->nb[2], src0->ne[2], stream);
4234
+ case GGML_TYPE_Q5_K:
4235
+ return reorder_qw_q5_k_moe(data_device, src0->nb[2], src0->ne[2], stream);
4236
+ case GGML_TYPE_Q6_K:
4237
+ return reorder_qw_q6_k_moe(data_device, src0->nb[2], src0->ne[2], stream);
4238
+ default:
4239
+ return false;
4240
+ }
4241
+ }
4242
+
4243
+ switch (src0->type) {
4244
+ case GGML_TYPE_Q4_0:
4245
+ return reorder_qw_q4_0(data_device, ncols, nrows, size, 0, stream);
4246
+ case GGML_TYPE_Q8_0:
4247
+ return reorder_qw_q8_0(data_device, ncols, nrows, size, 0, stream);
4248
+ case GGML_TYPE_Q3_K:
4249
+ return reorder_qw_q3_k(data_device, size, 0, stream);
3458
4250
  case GGML_TYPE_Q4_K:
3459
- reorder_qw_q4_k(data_device, size, 0, stream);
3460
- break;
4251
+ return reorder_qw_q4_k(data_device, size, 0, stream);
4252
+ case GGML_TYPE_Q5_K:
4253
+ return reorder_qw_q5_k(data_device, size, 0, stream);
3461
4254
  case GGML_TYPE_Q6_K:
3462
- reorder_qw_q6_k(data_device, size, 0, stream);
3463
- break;
4255
+ return reorder_qw_q6_k(data_device, size, 0, stream);
3464
4256
  default:
3465
- GGML_ABORT("reorder_qw() called with unsupported type");
3466
- break;
4257
+ return false;
3467
4258
  }
3468
4259
  }
3469
4260
 
3470
4261
  static bool should_reorder_tensor(ggml_backend_sycl_context& ctx, const ggml_tensor * dst) {
3471
- return !g_ggml_sycl_disable_optimize && //allow optimize, controlled by $GGML_SYCL_DISABLE_OPT
3472
- ctx.opt_feature.reorder && //allow this device due to good perf, skip the devices with bad perf.
3473
- dst->op == GGML_OP_MUL_MAT && //limit to some supported cases of Q4_0, to do for more cases.
3474
- dst->src[1]->ne[1]==1 && dst->src[1]->ne[2]==1 && dst->src[1]->ne[3]==1;
4262
+ return g_ggml_sycl_enable_optimize && //allow optimize, controlled by $GGML_SYCL_ENABLE_OPT
4263
+ ctx.opt_feature.reorder && //allow this device due to good perf, skip the devices with bad perf.
4264
+ dst->op == GGML_OP_MUL_MAT && //limit to some supported cases of Q4_0, to do for more cases.
4265
+ // ne[1] <= 8 so multi-column decode (spec / MTP verify) also bootstraps the reorder;
4266
+ // all reorderable types have a _switch_ncols kernel.
4267
+ dst->src[1]->ne[1] <= 8 && dst->src[1]->ne[2]==1 && dst->src[1]->ne[3]==1;
3475
4268
  }
3476
4269
 
3477
4270
  static void opt_for_reorder(ggml_backend_sycl_context * ctx, const ggml_tensor * src0, const ggml_tensor * /* src1 */,
@@ -3503,14 +4296,37 @@ static void opt_for_reorder(ggml_backend_sycl_context * ctx, const ggml_tensor *
3503
4296
  break;
3504
4297
  }
3505
4298
 
3506
- reorder_qw(src0, ctx->stream());
3507
- extra->optimized_feature.reorder = true; // Used to decode/dequan in next steps and avoid re-reordering
4299
+ if (reorder_qw(src0, ctx->stream())) {
4300
+ extra->optimized_feature.reorder = true; // Used to decode/dequan in next steps and avoid re-reordering
4301
+ }
4302
+ }
4303
+
4304
+ // Lazily reorder supported MoE expert weights once their fused path is used.
4305
+ static void opt_for_reorder_id(ggml_backend_sycl_context * ctx, const ggml_tensor * src0) {
4306
+ if (!g_ggml_sycl_enable_optimize || !ctx->opt_feature.reorder) {
4307
+ return;
4308
+ }
4309
+ if (src0->type != GGML_TYPE_Q4_K && src0->type != GGML_TYPE_Q5_K && src0->type != GGML_TYPE_Q6_K) {
4310
+ return;
4311
+ }
4312
+ ggml_tensor_extra_gpu * extra = static_cast<ggml_tensor_extra_gpu *>(src0->extra);
4313
+ if (!extra || extra->optimized_feature.reorder) {
4314
+ return;
4315
+ }
4316
+ if (reorder_qw(src0, ctx->stream())) {
4317
+ extra->optimized_feature.reorder = true;
4318
+ }
3508
4319
  }
3509
4320
 
3510
4321
 
3511
4322
  static bool can_use_dequantize_mul_mat_vec(const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) {
4323
+ // The F16/BF16 qk=1 kernel iterates with stride 2*DMMV_X, requiring ne[0] to be
4324
+ // a multiple of 2*DMMV_X. Quantized types use block-structured kernels that only
4325
+ // need ne[0] % DMMV_X == 0.
4326
+ const int64_t dmmv_x_required = (src0->type == GGML_TYPE_BF16 || src0->type == GGML_TYPE_F16) ?
4327
+ 2*GGML_SYCL_DMMV_X : GGML_SYCL_DMMV_X;
3512
4328
  return ggml_sycl_supports_dmmv(src0->type) && src1->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32 &&
3513
- src0->ne[0] % GGML_SYCL_DMMV_X == 0 && src1->ne[1] == 1;
4329
+ src0->ne[0] % dmmv_x_required == 0 && src1->ne[1] == 1;
3514
4330
  }
3515
4331
 
3516
4332
  static bool can_use_mul_mat_vec_q(const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) {
@@ -3560,9 +4376,16 @@ static void ggml_sycl_mul_mat(ggml_backend_sycl_context & ctx, const ggml_tensor
3560
4376
  // Dispatch becomes obscure with the reorder, MMVQ when the reorder optimization
3561
4377
  // is enabled takes precedence over DMMV, the current if-else implementation
3562
4378
  // requires disabling DMMV if both conditions are met
4379
+
3563
4380
  if (!g_ggml_sycl_prioritize_dmmv && ((should_reorder_tensor(ctx, dst) &&
3564
4381
  ggml_sycl_supports_reorder_mmvq(src0->type)))) {
3565
- use_dequantize_mul_mat_vec = use_dequantize_mul_mat_vec && !use_mul_mat_vec_q;
4382
+ // Arc770 get benefit with Q4_0 by skipping it.
4383
+ if (!(ggml_sycl_info().devices[ctx.device].hw_info.arch ==
4384
+ gpu_arch::intel_gpu_acm_g10 &&
4385
+ src0->type == GGML_TYPE_Q4_0)) {
4386
+ use_dequantize_mul_mat_vec =
4387
+ use_dequantize_mul_mat_vec && !use_mul_mat_vec_q;
4388
+ }
3566
4389
  }
3567
4390
 
3568
4391
  if (!split && src0->type == GGML_TYPE_F16 && ggml_is_permuted(src0) && ggml_is_permuted(src1) && src1->ne[1] == 1) {
@@ -3600,42 +4423,19 @@ static void ggml_sycl_mul_mat(ggml_backend_sycl_context & ctx, const ggml_tensor
3600
4423
  }
3601
4424
 
3602
4425
 
3603
- struct mmid_row_mapping {
3604
- int32_t i1;
3605
- int32_t i2;
3606
- };
3607
-
3608
4426
  __dpct_inline__ static void k_copy_src1_to_contiguous(
3609
4427
  const char *__restrict__ src1_original, char *__restrict__ src1_contiguous,
3610
- int *__restrict__ cur_src1_row, mmid_row_mapping *__restrict__ row_mapping,
3611
- const char *__restrict ids, int64_t i02, size_t ids_nb1, size_t ids_nb0,
4428
+ const mmid_row_mapping *__restrict__ row_mapping,
3612
4429
  int64_t ne11, int64_t ne10, size_t nb11, size_t nb12,
3613
- const sycl::nd_item<3> &item_ct1, int &src1_row) {
3614
- int32_t iid1 = item_ct1.get_group(2);
3615
- int32_t id = item_ct1.get_group(1);
4430
+ const sycl::nd_item<3> &item_ct1) {
4431
+ const int32_t src1_row = item_ct1.get_group(2);
3616
4432
 
3617
- const int32_t row_id_i = *(const int32_t *) (ids + iid1*ids_nb1 + id*ids_nb0);
3618
-
3619
- if (row_id_i != i02) {
3620
- return;
3621
- }
4433
+ const int32_t iid1 = row_mapping[src1_row].i2;
4434
+ const int32_t id = row_mapping[src1_row].i1;
3622
4435
 
3623
4436
  const int64_t i11 = id % ne11;
3624
4437
  const int64_t i12 = iid1;
3625
4438
 
3626
- if (item_ct1.get_local_id(2) == 0) {
3627
- src1_row =
3628
- dpct::atomic_fetch_add<sycl::access::address_space::generic_space>(
3629
- cur_src1_row, 1);
3630
- row_mapping[src1_row] = {id, iid1};
3631
- }
3632
- /*
3633
- DPCT1065:194: Consider replacing sycl::nd_item::barrier() with
3634
- sycl::nd_item::barrier(sycl::access::fence_space::local_space) for better
3635
- performance if there is no access to global memory.
3636
- */
3637
- item_ct1.barrier();
3638
-
3639
4439
  const float * src1_row_original = (const float *)(src1_original + i11*nb11 + i12*nb12);
3640
4440
  float * src1_row_contiguous = (float *)(src1_contiguous + src1_row*nb11);
3641
4441
 
@@ -3665,6 +4465,108 @@ __dpct_inline__ static void k_copy_dst_from_contiguous(
3665
4465
  }
3666
4466
  }
3667
4467
 
4468
+ // Fused MoE TG fast path. Returns false to fall back to the per-expert loop below.
4469
+ static bool ggml_sycl_mul_mat_id_mmvq_fused(
4470
+ ggml_backend_sycl_context & ctx, const ggml_tensor * src0,
4471
+ const ggml_tensor * src1, const ggml_tensor * ids, ggml_tensor * dst)
4472
+ {
4473
+ const int64_t ne10 = src1->ne[0];
4474
+ const int64_t ne11 = src1->ne[1];
4475
+ const int64_t ne12 = src1->ne[2];
4476
+ if (ne12 != 1) return false;
4477
+ if (src1->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_F32) return false;
4478
+ if (ne10 != src0->ne[0] || ne10 % QK8_1 != 0) return false;
4479
+ if (!ggml_is_contiguous(src1)) return false;
4480
+
4481
+ const int64_t n_ids_per_group = ids->ne[0];
4482
+ if (ids->ne[1] != 1) return false;
4483
+ if (ne11 != 1 && ne11 != n_ids_per_group) return false;
4484
+
4485
+ const queue_ptr stream = ctx.stream();
4486
+ const int src1_padded_cols = GGML_PAD((int) ne10, MATRIX_ROW_PADDING);
4487
+ const int n_experts_used = (int) n_ids_per_group;
4488
+ const int nrows = (int) src0->ne[1];
4489
+
4490
+ // Lazily reorder the (Q4_K) expert weights into a per-expert SoA layout, then run the reorder
4491
+ // GEMV. Placed after the bail checks so a non-dispatchable op does not pay the reorder cost.
4492
+ opt_for_reorder_id(&ctx, src0);
4493
+ const ggml_tensor_extra_gpu * src0_extra =
4494
+ static_cast<const ggml_tensor_extra_gpu *>(src0->extra);
4495
+ const bool use_reorder = src0_extra && src0_extra->optimized_feature.reorder;
4496
+
4497
+ ggml_sycl_pool_alloc<char> src1_q8_alloc(ctx.pool(),
4498
+ (size_t) ne11 * src1_padded_cols * sizeof(block_q8_1) / QK8_1);
4499
+ char * src1_ddq = src1_q8_alloc.get();
4500
+ if (use_reorder) {
4501
+ quantize_row_q8_1_sycl<quantize_and_reorder_q8_1_soa>(
4502
+ (const float *) src1->data, src1_ddq, (int) ne10, (int) ne11,
4503
+ src1_padded_cols, stream);
4504
+ } else {
4505
+ quantize_row_q8_1_sycl<quantize_q8_1>(
4506
+ (const float *) src1->data, src1_ddq, (int) ne10, (int) ne11,
4507
+ src1_padded_cols, stream);
4508
+ }
4509
+
4510
+ const size_t bytes_per_qrow = (size_t) src1_padded_cols * sizeof(block_q8_1) / QK8_1;
4511
+ const size_t src1_row_stride = (ne11 == 1) ? 0 : bytes_per_qrow;
4512
+
4513
+ if (use_reorder) {
4514
+ return ggml_sycl_mul_mat_vec_q_id_reorder(
4515
+ src0->type, src0->data, src1_ddq, (const int32_t *) ids->data,
4516
+ (float *) dst->data, (int) ne10, nrows, n_experts_used,
4517
+ /*expert_weight_stride=*/ src0->nb[2],
4518
+ /*dst_row_stride=*/ dst->nb[1],
4519
+ src1_row_stride, stream);
4520
+ }
4521
+ return ggml_sycl_mul_mat_vec_q_id(
4522
+ src0->type, src0->data, src1_ddq, (const int32_t *) ids->data,
4523
+ (float *) dst->data, (int) ne10, nrows, n_experts_used,
4524
+ /*expert_weight_stride=*/ src0->nb[2],
4525
+ /*dst_row_stride=*/ dst->nb[1],
4526
+ src1_row_stride, stream);
4527
+ }
4528
+
4529
+ // counting sort of the routed rows by expert id (row_id_i, as chosen by the router):
4530
+ // builds a projection of a memory layout where each expert's slice is contiguous
4531
+ static void mmid_counting_sort_rows(
4532
+ const ggml_tensor * ids, const char * ids_host,
4533
+ int64_t n_ids, int64_t n_as, int64_t n_routed_rows,
4534
+ std::vector<int64_t> & expert_counts,
4535
+ std::vector<int64_t> & expert_row_offsets,
4536
+ std::vector<mmid_row_mapping> & routed_row_src) {
4537
+
4538
+ // frequencies: how many routed rows each expert "owns"
4539
+ expert_counts.assign(n_as, 0);
4540
+ for (int64_t iid1 = 0; iid1 < ids->ne[1]; iid1++) {
4541
+ for (int64_t id = 0; id < n_ids; id++) {
4542
+ const int32_t row_id_i = *(const int32_t *) (ids_host + iid1*ids->nb[1] + id*ids->nb[0]);
4543
+ GGML_ASSERT(row_id_i >= 0 && row_id_i < n_as);
4544
+ expert_counts[row_id_i]++;
4545
+ }
4546
+ }
4547
+
4548
+ // where each expert's slice starts (row indices) and the previous ends
4549
+ expert_row_offsets.assign(n_as + 1, 0);
4550
+ for (int64_t i02 = 0; i02 < n_as; i02++) {
4551
+ expert_row_offsets[i02 + 1] = expert_row_offsets[i02] + expert_counts[i02];
4552
+ }
4553
+
4554
+ std::vector<int64_t> expert_row_next = expert_row_offsets;
4555
+ routed_row_src.resize(n_routed_rows);
4556
+ for (int64_t iid1 = 0; iid1 < ids->ne[1]; iid1++) {
4557
+ for (int64_t id = 0; id < n_ids; id++) {
4558
+ const int32_t row_id_i = *(const int32_t *) (ids_host + iid1*ids->nb[1] + id*ids->nb[0]);
4559
+ GGML_ASSERT(row_id_i >= 0 && row_id_i < n_as);
4560
+
4561
+ // find and validate the next free row for a given expert (row_id_i)
4562
+ const int64_t routed_row = expert_row_next[row_id_i]++;
4563
+ GGML_ASSERT(routed_row >= expert_row_offsets[row_id_i]);
4564
+ GGML_ASSERT(routed_row < expert_row_offsets[row_id_i + 1]);
4565
+ routed_row_src[routed_row] = {(int32_t) id, (int32_t) iid1};
4566
+ }
4567
+ }
4568
+ }
4569
+
3668
4570
  static void ggml_sycl_mul_mat_id(ggml_backend_sycl_context & ctx,
3669
4571
  ggml_tensor *dst) try {
3670
4572
  scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/3);
@@ -3680,11 +4582,19 @@ static void ggml_sycl_mul_mat_id(ggml_backend_sycl_context & ctx,
3680
4582
  const int64_t n_as = ne02;
3681
4583
  const int64_t n_ids = ids->ne[0];
3682
4584
 
4585
+ if (ne12 == 1) {
4586
+ if (ggml_sycl_mul_mat_id_mmvq_fused(ctx, src0, src1, ids, dst)) {
4587
+ return;
4588
+ }
4589
+ }
4590
+
3683
4591
  std::vector<char> ids_host(ggml_nbytes(ids));
3684
4592
  const char * ids_dev = (const char *) ids->data;
3685
4593
 
3686
4594
  SYCL_CHECK(CHECK_TRY_ERROR(
3687
4595
  stream->memcpy(ids_host.data(), ids_dev, ggml_nbytes(ids))));
4596
+
4597
+ // also ensures ctx.mmid_row_mapping_host is drained before we use it again
3688
4598
  SYCL_CHECK(CHECK_TRY_ERROR(stream->wait()));
3689
4599
 
3690
4600
  ggml_tensor src0_row = *src0;
@@ -3730,105 +4640,98 @@ static void ggml_sycl_mul_mat_id(ggml_backend_sycl_context & ctx,
3730
4640
  }
3731
4641
  }
3732
4642
  } else {
3733
- ggml_sycl_pool_alloc<char> src1_contiguous(ctx.pool(), sizeof(float)*ggml_nelements(src1));
3734
- ggml_sycl_pool_alloc<char> dst_contiguous(ctx.pool(), sizeof(float)*ggml_nelements(dst));
4643
+ const int64_t n_routed_rows = ids->ne[1] * n_ids;
4644
+ ggml_sycl_pool_alloc<char> src1_contiguous(ctx.pool(), sizeof(float)*n_routed_rows*ne10);
4645
+ ggml_sycl_pool_alloc<char> dst_contiguous(ctx.pool(), sizeof(float)*n_routed_rows*ne0);
3735
4646
 
3736
4647
  src1_row.data = src1_contiguous.get();
3737
4648
  dst_row.data = dst_contiguous.get();
3738
4649
 
3739
- for (int64_t i02 = 0; i02 < n_as; i02++) {
3740
- int64_t num_src1_rows = 0;
3741
- for (int64_t iid1 = 0; iid1 < ids->ne[1]; iid1++) {
3742
- for (int64_t id = 0; id < n_ids; id++) {
3743
- const int32_t row_id_i = *(const int32_t *) (ids_host.data() + iid1*ids->nb[1] + id*ids->nb[0]);
4650
+ // how many "owned" routed rows to pass to each expert
4651
+ std::vector<int64_t> expert_row_counts;
4652
+ // where each expert's slice starts and the previous ends (row indices, right-exclusive)
4653
+ std::vector<int64_t> expert_row_offsets;
4654
+ // the sources (slot/token pairs) of contiguous rows to guide k_copy_src1_to_contiguous
4655
+ std::vector<mmid_row_mapping> & routed_row_src = ctx.mmid_row_mapping_host;
3744
4656
 
3745
- GGML_ASSERT(row_id_i >= 0 && row_id_i < n_as);
4657
+ mmid_counting_sort_rows(ids, ids_host.data(), n_ids, n_as, n_routed_rows,
4658
+ expert_row_counts, expert_row_offsets, routed_row_src);
3746
4659
 
3747
- if (row_id_i != i02) {
3748
- continue;
3749
- }
4660
+ ggml_sycl_pool_alloc<mmid_row_mapping> dev_row_mapping(ctx.pool(), n_routed_rows);
4661
+ SYCL_CHECK(CHECK_TRY_ERROR(
4662
+ stream->memcpy(dev_row_mapping.get(), routed_row_src.data(), n_routed_rows*sizeof(mmid_row_mapping))));
3750
4663
 
3751
- num_src1_rows++;
3752
- }
3753
- }
4664
+ const unsigned int max_work_group_size = ggml_sycl_info().max_work_group_sizes[ctx.device];
4665
+ assert(max_work_group_size % (WARP_SIZE * WARP_SIZE) == 0);
4666
+
4667
+ {
4668
+ sycl::range<3> block_dims(1, 1, std::min((unsigned int)ne10, max_work_group_size));
4669
+ sycl::range<3> grid_dims(1, 1, n_routed_rows);
4670
+ stream->submit([&](sycl::handler &cgh) {
4671
+ char *__restrict src1_contiguous_get =
4672
+ src1_contiguous.get();
4673
+ mmid_row_mapping *__restrict dev_row_mapping_get =
4674
+ dev_row_mapping.get();
4675
+
4676
+ cgh.parallel_for(
4677
+ sycl::nd_range<3>(grid_dims * block_dims, block_dims),
4678
+ [=](sycl::nd_item<3> item_ct1) {
4679
+ k_copy_src1_to_contiguous(
4680
+ src1_original, src1_contiguous_get,
4681
+ dev_row_mapping_get,
4682
+ ne11, ne10, nb11, nb12,
4683
+ item_ct1);
4684
+ });
4685
+ });
4686
+ }
4687
+
4688
+ for (int64_t i02 = 0; i02 < n_as; i02++) {
4689
+ const int64_t num_src1_rows = expert_row_counts[i02];
3754
4690
 
3755
4691
  if (num_src1_rows == 0) {
3756
4692
  continue;
3757
4693
  }
3758
4694
 
3759
-
3760
- ggml_sycl_pool_alloc<int> dev_cur_src1_row(ctx.pool(), 1);
3761
- ggml_sycl_pool_alloc<mmid_row_mapping> dev_row_mapping(ctx.pool(), num_src1_rows);
3762
- SYCL_CHECK(CHECK_TRY_ERROR(
3763
- stream->memset(dev_cur_src1_row.get(), 0, sizeof(int))));
3764
-
3765
- const unsigned int max_work_group_size = ggml_sycl_info().max_work_group_sizes[ctx.device];
3766
- assert(max_work_group_size % (WARP_SIZE * WARP_SIZE) == 0);
3767
-
3768
- {
3769
- sycl::range<3> block_dims(1, 1, std::min((unsigned int)ne10, max_work_group_size));
3770
- sycl::range<3> grid_dims(1, n_ids, ids->ne[1]);
3771
- stream->submit([&](sycl::handler &cgh) {
3772
- sycl::local_accessor<int, 0> src1_row_acc(cgh);
3773
-
3774
- char *__restrict src1_contiguous_get =
3775
- src1_contiguous.get();
3776
- int *__restrict dev_cur_src1_row_get =
3777
- dev_cur_src1_row.get();
3778
- mmid_row_mapping *__restrict dev_row_mapping_get =
3779
- dev_row_mapping.get();
3780
- size_t ids_nb_ct6 = ids->nb[1];
3781
- size_t ids_nb_ct7 = ids->nb[0];
3782
-
3783
- cgh.parallel_for(
3784
- sycl::nd_range<3>(grid_dims * block_dims, block_dims),
3785
- [=](sycl::nd_item<3> item_ct1) {
3786
- k_copy_src1_to_contiguous(
3787
- src1_original, src1_contiguous_get,
3788
- dev_cur_src1_row_get,
3789
- dev_row_mapping_get, ids_dev, i02,
3790
- ids_nb_ct6, ids_nb_ct7, ne11, ne10, nb11, nb12,
3791
- item_ct1, src1_row_acc);
3792
- });
3793
- });
3794
- }
4695
+ const int64_t expert_row_offset = expert_row_offsets[i02];
3795
4696
 
3796
4697
  src0_row.data = src0_original + i02*nb02;
3797
4698
 
3798
4699
  GGML_ASSERT(nb11 == sizeof(float)*ne10);
3799
4700
  GGML_ASSERT(nb1 == sizeof(float)*ne0);
4701
+ src1_row.data = src1_contiguous.get() + expert_row_offset*nb11;
3800
4702
  src1_row.ne[1] = num_src1_rows;
3801
4703
 
3802
4704
  src1_row.nb[1] = nb11;
3803
4705
  src1_row.nb[2] = num_src1_rows*nb11;
3804
4706
  src1_row.nb[3] = num_src1_rows*nb11;
3805
4707
 
4708
+ dst_row.data = dst_contiguous.get() + expert_row_offset*nb1;
3806
4709
  dst_row.ne[1] = num_src1_rows;
3807
4710
  dst_row.nb[1] = nb1;
3808
4711
  dst_row.nb[2] = num_src1_rows*nb1;
3809
4712
  dst_row.nb[3] = num_src1_rows*nb1;
3810
4713
 
3811
4714
  ggml_sycl_mul_mat(ctx, &src0_row, &src1_row, &dst_row);
4715
+ }
3812
4716
 
3813
- {
3814
- sycl::range<3> block_dims(1, 1, std::min((unsigned int)ne0, max_work_group_size));
3815
- sycl::range<3> grid_dims(1, 1, num_src1_rows);
3816
- stream->submit([&](sycl::handler &cgh) {
3817
- const char *__restrict dst_contiguous_get =
3818
- dst_contiguous.get();
3819
- const mmid_row_mapping *__restrict dev_row_mapping_get =
3820
- dev_row_mapping.get();
3821
-
3822
- cgh.parallel_for(
3823
- sycl::nd_range<3>(grid_dims * block_dims, block_dims),
3824
- [=](sycl::nd_item<3> item_ct1) {
3825
- k_copy_dst_from_contiguous(dst_original,
3826
- dst_contiguous_get,
3827
- dev_row_mapping_get,
3828
- ne0, nb1, nb2, item_ct1);
3829
- });
3830
- });
3831
- }
4717
+ {
4718
+ sycl::range<3> block_dims(1, 1, std::min((unsigned int)ne0, max_work_group_size));
4719
+ sycl::range<3> grid_dims(1, 1, n_routed_rows);
4720
+ stream->submit([&](sycl::handler &cgh) {
4721
+ const char *__restrict dst_contiguous_get =
4722
+ dst_contiguous.get();
4723
+ const mmid_row_mapping *__restrict dev_row_mapping_get =
4724
+ dev_row_mapping.get();
4725
+
4726
+ cgh.parallel_for(
4727
+ sycl::nd_range<3>(grid_dims * block_dims, block_dims),
4728
+ [=](sycl::nd_item<3> item_ct1) {
4729
+ k_copy_dst_from_contiguous(dst_original,
4730
+ dst_contiguous_get,
4731
+ dev_row_mapping_get,
4732
+ ne0, nb1, nb2, item_ct1);
4733
+ });
4734
+ });
3832
4735
  }
3833
4736
  }
3834
4737
  }
@@ -3853,11 +4756,31 @@ static void ggml_sycl_pool2d(ggml_backend_sycl_context & ctx, ggml_tensor * dst)
3853
4756
  ggml_sycl_op_pool2d(ctx, dst);
3854
4757
  }
3855
4758
 
4759
+ static void ggml_sycl_pool1d(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
4760
+ scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/1);
4761
+ ggml_sycl_op_pool1d(ctx, dst);
4762
+ }
4763
+
3856
4764
  static void ggml_sycl_im2col(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
3857
4765
  scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/2);
3858
4766
  ggml_sycl_op_im2col(ctx, dst);
3859
4767
  }
3860
4768
 
4769
+ static void ggml_sycl_im2col_3d(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
4770
+ scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/2);
4771
+ ggml_sycl_op_im2col_3d(ctx, dst);
4772
+ }
4773
+
4774
+ static void ggml_sycl_col2im_1d(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
4775
+ scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/1);
4776
+ ggml_sycl_op_col2im_1d(ctx, dst);
4777
+ }
4778
+
4779
+ static void ggml_sycl_conv_3d(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
4780
+ scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/2);
4781
+ ggml_sycl_op_conv_3d(ctx, dst);
4782
+ }
4783
+
3861
4784
  static void ggml_sycl_sum(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
3862
4785
  scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/1);
3863
4786
  GGML_ASSERT(ggml_is_contiguous(dst->src[0]));
@@ -3921,9 +4844,21 @@ static bool ggml_sycl_compute_forward(ggml_backend_sycl_context & ctx, struct gg
3921
4844
  case GGML_OP_ARGMAX:
3922
4845
  ggml_sycl_argmax(ctx, dst);
3923
4846
  break;
4847
+ case GGML_OP_CONV_2D:
4848
+ ggml_sycl_op_conv2d(ctx, dst);
4849
+ break;
4850
+ case GGML_OP_CONV_2D_DW:
4851
+ ggml_sycl_op_conv2d_dw(ctx, dst);
4852
+ break;
4853
+ case GGML_OP_CONV_3D:
4854
+ ggml_sycl_conv_3d(ctx, dst);
4855
+ break;
3924
4856
  case GGML_OP_CONV_TRANSPOSE_1D:
3925
4857
  ggml_sycl_op_conv_transpose_1d(ctx, dst);
3926
4858
  break;
4859
+ case GGML_OP_CONV_TRANSPOSE_2D:
4860
+ ggml_sycl_op_conv2d_transpose(ctx, dst);
4861
+ break;
3927
4862
  case GGML_OP_REPEAT:
3928
4863
  ggml_sycl_repeat(ctx, dst);
3929
4864
  break;
@@ -4005,6 +4940,9 @@ static bool ggml_sycl_compute_forward(ggml_backend_sycl_context & ctx, struct gg
4005
4940
  case GGML_UNARY_OP_EXP:
4006
4941
  ggml_sycl_exp(ctx, dst);
4007
4942
  break;
4943
+ case GGML_UNARY_OP_EXPM1:
4944
+ ggml_sycl_expm1(ctx, dst);
4945
+ break;
4008
4946
  case GGML_UNARY_OP_SOFTPLUS:
4009
4947
  ggml_sycl_softplus(ctx, dst);
4010
4948
  break;
@@ -4146,6 +5084,12 @@ static bool ggml_sycl_compute_forward(ggml_backend_sycl_context & ctx, struct gg
4146
5084
  case GGML_OP_SOFT_MAX_BACK:
4147
5085
  ggml_sycl_op_soft_max_back(ctx, dst);
4148
5086
  break;
5087
+ case GGML_OP_CROSS_ENTROPY_LOSS:
5088
+ ggml_sycl_cross_entropy_loss(ctx, dst);
5089
+ break;
5090
+ case GGML_OP_CROSS_ENTROPY_LOSS_BACK:
5091
+ ggml_sycl_cross_entropy_loss_back(ctx, dst);
5092
+ break;
4149
5093
  case GGML_OP_ROPE:
4150
5094
  ggml_sycl_rope(ctx, dst);
4151
5095
  break;
@@ -4155,9 +5099,18 @@ static bool ggml_sycl_compute_forward(ggml_backend_sycl_context & ctx, struct gg
4155
5099
  case GGML_OP_IM2COL:
4156
5100
  ggml_sycl_im2col(ctx, dst);
4157
5101
  break;
5102
+ case GGML_OP_IM2COL_3D:
5103
+ ggml_sycl_im2col_3d(ctx, dst);
5104
+ break;
5105
+ case GGML_OP_COL2IM_1D:
5106
+ ggml_sycl_col2im_1d(ctx, dst);
5107
+ break;
4158
5108
  case GGML_OP_POOL_2D:
4159
5109
  ggml_sycl_pool2d(ctx, dst);
4160
5110
  break;
5111
+ case GGML_OP_POOL_1D:
5112
+ ggml_sycl_pool1d(ctx, dst);
5113
+ break;
4161
5114
  case GGML_OP_SUM:
4162
5115
  ggml_sycl_sum(ctx, dst);
4163
5116
  break;
@@ -4191,6 +5144,21 @@ static bool ggml_sycl_compute_forward(ggml_backend_sycl_context & ctx, struct gg
4191
5144
  case GGML_OP_SSM_CONV:
4192
5145
  ggml_sycl_ssm_conv(ctx, dst);
4193
5146
  break;
5147
+ case GGML_OP_SSM_SCAN:
5148
+ ggml_sycl_ssm_scan(ctx, dst);
5149
+ break;
5150
+ case GGML_OP_FILL:
5151
+ ggml_sycl_fill(ctx, dst);
5152
+ break;
5153
+ case GGML_OP_CUMSUM:
5154
+ ggml_sycl_cumsum(ctx, dst);
5155
+ break;
5156
+ case GGML_OP_DIAG:
5157
+ ggml_sycl_diag(ctx, dst);
5158
+ break;
5159
+ case GGML_OP_SOLVE_TRI:
5160
+ ggml_sycl_solve_tri(ctx, dst);
5161
+ break;
4194
5162
  case GGML_OP_ROLL:
4195
5163
  ggml_sycl_roll(ctx, dst);
4196
5164
  break;
@@ -4417,7 +5385,10 @@ static ggml_status ggml_backend_sycl_graph_compute(ggml_backend_t backend, ggml_
4417
5385
  auto * sycl_ctx = static_cast<ggml_backend_sycl_context *>(backend->context);
4418
5386
 
4419
5387
  #ifdef GGML_SYCL_GRAPH
4420
- bool use_sycl_graph = !g_ggml_sycl_disable_graph && check_graph_compatibility(cgraph);
5388
+ bool use_sycl_graph = false;
5389
+ if (g_ggml_sycl_enable_graph) {
5390
+ use_sycl_graph = check_graph_compatibility(cgraph);
5391
+ }
4421
5392
  if (use_sycl_graph) {
4422
5393
  const bool graph_support = dpct::get_device(sycl_ctx->device).has(sycl::aspect::ext_oneapi_limited_graph);
4423
5394
  if (!graph_support) {
@@ -4497,6 +5468,8 @@ static ggml_backend_i ggml_backend_sycl_interface = {
4497
5468
  /* .free = */ ggml_backend_sycl_free,
4498
5469
  /* .set_tensor_async = */ ggml_backend_sycl_set_tensor_async,
4499
5470
  /* .get_tensor_async = */ ggml_backend_sycl_get_tensor_async,
5471
+ /* .set_tensor_2d_async = */ NULL,
5472
+ /* .get_tensor_2d_async = */ NULL,
4500
5473
  /* .cpy_tensor_async = */ NULL, // ggml_backend_sycl_cpy_tensor_async,
4501
5474
  // // TODO: update for the new
4502
5475
  // interface
@@ -4601,7 +5574,7 @@ static ggml_backend_buffer_t ggml_backend_sycl_device_buffer_from_host_ptr(ggml_
4601
5574
  return nullptr;
4602
5575
  }
4603
5576
 
4604
- static bool ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const ggml_tensor * op) {
5577
+ static bool do_ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const ggml_tensor * op) {
4605
5578
  ggml_backend_sycl_device_context *sycl_ctx =
4606
5579
  (ggml_backend_sycl_device_context *)dev->context;
4607
5580
  int device = sycl_ctx->device;
@@ -4615,6 +5588,10 @@ static bool ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const g
4615
5588
  }
4616
5589
  return false;
4617
5590
  }
5591
+ case GGML_OP_CONV_2D:
5592
+ case GGML_OP_CONV_2D_DW:
5593
+ case GGML_OP_CONV_TRANSPOSE_2D:
5594
+ return true;
4618
5595
  case GGML_OP_UNARY:
4619
5596
  switch (ggml_get_unary_op(op)) {
4620
5597
  case GGML_UNARY_OP_SGN:
@@ -4631,6 +5608,7 @@ static bool ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const g
4631
5608
  case GGML_UNARY_OP_GELU_QUICK:
4632
5609
  case GGML_UNARY_OP_GELU_ERF:
4633
5610
  case GGML_UNARY_OP_EXP:
5611
+ case GGML_UNARY_OP_EXPM1:
4634
5612
  case GGML_UNARY_OP_SOFTPLUS:
4635
5613
  case GGML_UNARY_OP_ELU:
4636
5614
  case GGML_UNARY_OP_CEIL:
@@ -4638,11 +5616,7 @@ static bool ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const g
4638
5616
  case GGML_UNARY_OP_FLOOR:
4639
5617
  case GGML_UNARY_OP_ROUND:
4640
5618
  case GGML_UNARY_OP_TRUNC:
4641
- #if defined (GGML_SYCL_F16)
4642
- return ggml_is_contiguous(op->src[0]) && (op->type == op->src[0]->type);
4643
- #else
4644
- return ggml_is_contiguous(op->src[0]) && (op->src[0]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32) && (op->type == op->src[0]->type);
4645
- #endif
5619
+ return true;
4646
5620
  default:
4647
5621
  return false;
4648
5622
  }
@@ -4668,22 +5642,8 @@ static bool ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const g
4668
5642
  if (a->ne[3] != b->ne[3]) {
4669
5643
  return false;
4670
5644
  }
4671
- ggml_type a_type = a->type;
4672
- if (a_type == GGML_TYPE_IQ4_NL || a_type == GGML_TYPE_IQ4_XS ||
4673
- a_type == GGML_TYPE_IQ3_XXS || a_type == GGML_TYPE_IQ3_S ||
4674
- a_type == GGML_TYPE_IQ2_XXS || a_type == GGML_TYPE_IQ2_XS || a_type == GGML_TYPE_IQ2_S ||
4675
- a_type == GGML_TYPE_IQ1_S || a_type == GGML_TYPE_IQ1_M
4676
- ) {
4677
- if (b->ne[1] == 1 && ggml_nrows(b) > 1) {
4678
- return false;
4679
- }
4680
- }
5645
+
4681
5646
  ggml_type src0_type = op->src[0]->type;
4682
- if (src0_type == GGML_TYPE_BF16 ) {
4683
- // TODO: support GGML_TYPE_BF16
4684
- // FIXME: keep a list of supported types to avoid breaking the backend when a new type is added
4685
- return false;
4686
- }
4687
5647
 
4688
5648
  // TODO: The configuration below needs more work to be supported with oneDNN
4689
5649
  if (ggml_is_permuted(a) && !ggml_is_contiguous(a) &&
@@ -4699,16 +5659,39 @@ static bool ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const g
4699
5659
  return true;
4700
5660
  }
4701
5661
  case GGML_OP_OUT_PROD:
4702
- return op->type == GGML_TYPE_F32 && op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && op->ne[2] == 1 && op->ne[3] == 1;
5662
+ return op->type == GGML_TYPE_F32 &&
5663
+ (op->src[0]->type == GGML_TYPE_F32 ||
5664
+ (op->src[0]->type == GGML_TYPE_Q1_0 && op->src[0]->ne[2] == op->src[1]->ne[2] &&
5665
+ op->src[0]->ne[3] == op->src[1]->ne[3])) &&
5666
+ op->src[1]->type == GGML_TYPE_F32;
4703
5667
  case GGML_OP_GET_ROWS:
4704
5668
  {
4705
5669
  switch (op->src[0]->type) {
5670
+ case GGML_TYPE_I32:
4706
5671
  case GGML_TYPE_F16:
5672
+ case GGML_TYPE_BF16:
4707
5673
  case GGML_TYPE_F32:
5674
+ case GGML_TYPE_Q1_0:
5675
+ case GGML_TYPE_MXFP4:
5676
+ case GGML_TYPE_NVFP4:
5677
+ case GGML_TYPE_IQ2_XXS:
5678
+ case GGML_TYPE_IQ2_XS:
5679
+ case GGML_TYPE_IQ2_S:
5680
+ case GGML_TYPE_IQ3_XXS:
5681
+ case GGML_TYPE_IQ1_S:
5682
+ case GGML_TYPE_IQ1_M:
5683
+ case GGML_TYPE_IQ3_S:
5684
+ case GGML_TYPE_IQ4_NL:
5685
+ case GGML_TYPE_IQ4_XS:
5686
+ case GGML_TYPE_Q2_K:
5687
+ case GGML_TYPE_Q3_K:
4708
5688
  case GGML_TYPE_Q4_0:
4709
5689
  case GGML_TYPE_Q4_1:
5690
+ case GGML_TYPE_Q4_K:
4710
5691
  case GGML_TYPE_Q5_0:
4711
5692
  case GGML_TYPE_Q5_1:
5693
+ case GGML_TYPE_Q5_K:
5694
+ case GGML_TYPE_Q6_K:
4712
5695
  case GGML_TYPE_Q8_0:
4713
5696
  return true;
4714
5697
  default:
@@ -4723,80 +5706,114 @@ static bool ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const g
4723
5706
 
4724
5707
  case GGML_OP_SET_ROWS:
4725
5708
  {
4726
- return ((op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_F16 || op->type == GGML_TYPE_BF16 ||
5709
+
5710
+ auto res = ((op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_F16 || op->type == GGML_TYPE_BF16 ||
4727
5711
  op->type == GGML_TYPE_Q8_0 || op->type == GGML_TYPE_Q5_1 || op->type == GGML_TYPE_Q5_0 ||
4728
- op->type == GGML_TYPE_Q4_1 || op->type == GGML_TYPE_Q4_0 || op->type == GGML_TYPE_IQ4_NL) &&
5712
+ op->type == GGML_TYPE_Q1_0 ||
5713
+ op->type == GGML_TYPE_Q4_1 || op->type == GGML_TYPE_Q4_0 || op->type == GGML_TYPE_IQ4_NL ||
5714
+ op->type == GGML_TYPE_MXFP4 || op->type == GGML_TYPE_NVFP4) &&
5715
+ op->src[0]->type == GGML_TYPE_F32 &&
4729
5716
  (op->src[1]->type == GGML_TYPE_I64 || op->src[1]->type == GGML_TYPE_I32));
5717
+ return res;
4730
5718
  }
4731
5719
  break;
4732
5720
  case GGML_OP_CPY:
4733
5721
  {
4734
5722
  ggml_type src0_type = op->src[0]->type;
4735
5723
  ggml_type src1_type = op->src[1]->type;
4736
- if (src0_type == src1_type && (ggml_is_contiguous(op->src[0]) && ggml_is_contiguous(op->src[1])) && src0_type != GGML_TYPE_BF16) {
4737
- return true;
4738
- }
4739
- if (src0_type == GGML_TYPE_F32 && src1_type == GGML_TYPE_F32) {
4740
- return true;
4741
- }
4742
- if (src0_type == GGML_TYPE_F32 && src1_type == GGML_TYPE_F16) {
4743
- return true;
4744
- }
4745
- if (src0_type == GGML_TYPE_F32 && src1_type == GGML_TYPE_Q8_0) {
4746
- return true;
4747
- }
4748
- if (src0_type == GGML_TYPE_F32 && src1_type == GGML_TYPE_Q4_0) {
4749
- return true;
4750
- }
4751
- if (src0_type == GGML_TYPE_F32 && src1_type == GGML_TYPE_Q4_1) {
4752
- return true;
4753
- }
4754
- if (src0_type == GGML_TYPE_F16 && src1_type == GGML_TYPE_F16) {
4755
- return true;
4756
- }
4757
- if (src0_type == GGML_TYPE_F16 && src1_type == GGML_TYPE_F32) {
4758
- return true;
4759
- }
4760
- if (src0_type == GGML_TYPE_Q8_0 && src1_type == GGML_TYPE_F32) {
4761
- return true;
4762
- }
4763
- if (src0_type == GGML_TYPE_Q4_0 && src1_type == GGML_TYPE_F32) {
4764
- return true;
4765
- }
4766
- if (src0_type == GGML_TYPE_Q4_1 && src1_type == GGML_TYPE_F32) {
4767
- return true;
4768
- }
4769
- if (src0_type == GGML_TYPE_F32 && src1_type == GGML_TYPE_Q5_0) {
4770
- return true;
4771
- }
4772
- if (src0_type == GGML_TYPE_Q5_0 && src1_type == GGML_TYPE_F32) {
4773
- return true;
4774
- }
4775
- if (src0_type == GGML_TYPE_F32 && src1_type == GGML_TYPE_Q5_1) {
4776
- return true;
4777
- }
4778
- if (src0_type == GGML_TYPE_Q5_1 && src1_type == GGML_TYPE_F32) {
4779
- return true;
4780
- }
4781
- if (src0_type == GGML_TYPE_F32 && src1_type == GGML_TYPE_IQ4_NL) {
4782
- return true;
4783
- }
4784
- if(src0_type == GGML_TYPE_Q8_0 && src1_type == GGML_TYPE_Q8_0) {
4785
- return true;
5724
+
5725
+ if (src0_type == GGML_TYPE_F16) {
5726
+ if (src1_type == GGML_TYPE_Q2_K ||
5727
+ src1_type == GGML_TYPE_Q3_K ||
5728
+ src1_type == GGML_TYPE_Q4_K ||
5729
+ src1_type == GGML_TYPE_Q5_K ||
5730
+ src1_type == GGML_TYPE_Q6_K ||
5731
+ src1_type == GGML_TYPE_IQ2_XXS ||
5732
+ src1_type == GGML_TYPE_IQ2_XS ||
5733
+ src1_type == GGML_TYPE_IQ2_S ||
5734
+ src1_type == GGML_TYPE_IQ3_XXS ||
5735
+ src1_type == GGML_TYPE_IQ1_S ||
5736
+ src1_type == GGML_TYPE_IQ1_M ||
5737
+ src1_type == GGML_TYPE_IQ3_S ||
5738
+ src1_type == GGML_TYPE_IQ4_XS) {
5739
+ return false;
5740
+ }
4786
5741
  }
4787
- if(src0_type == GGML_TYPE_Q5_0 && src1_type == GGML_TYPE_Q5_0) {
4788
- return true;
5742
+
5743
+ if (src0_type == GGML_TYPE_BF16) {
5744
+ if (src1_type == GGML_TYPE_Q4_0 || //big error in ut
5745
+ src1_type == GGML_TYPE_Q4_1 || //big error in ut
5746
+ src1_type == GGML_TYPE_Q8_0 || //big error in ut
5747
+ src1_type == GGML_TYPE_Q2_K ||
5748
+ src1_type == GGML_TYPE_Q3_K ||
5749
+ src1_type == GGML_TYPE_Q4_K ||
5750
+ src1_type == GGML_TYPE_Q5_K ||
5751
+ src1_type == GGML_TYPE_Q6_K ||
5752
+ src1_type == GGML_TYPE_IQ2_XXS ||
5753
+ src1_type == GGML_TYPE_IQ2_XS ||
5754
+ src1_type == GGML_TYPE_IQ2_S ||
5755
+ src1_type == GGML_TYPE_IQ3_XXS ||
5756
+ src1_type == GGML_TYPE_IQ1_S ||
5757
+ src1_type == GGML_TYPE_IQ1_M ||
5758
+ src1_type == GGML_TYPE_IQ3_S ||
5759
+ src1_type == GGML_TYPE_IQ4_XS) {
5760
+ return false;
5761
+ }
4789
5762
  }
4790
- if(src0_type == GGML_TYPE_Q5_1 && src1_type == GGML_TYPE_Q5_1) {
4791
- return true;
5763
+
5764
+ if (src0_type == GGML_TYPE_F32) {
5765
+ if (src1_type == GGML_TYPE_Q2_K ||
5766
+ src1_type == GGML_TYPE_Q3_K ||
5767
+ src1_type == GGML_TYPE_Q4_K ||
5768
+ src1_type == GGML_TYPE_Q5_K ||
5769
+ src1_type == GGML_TYPE_Q6_K ||
5770
+ src1_type == GGML_TYPE_IQ2_XXS ||
5771
+ src1_type == GGML_TYPE_IQ2_XS ||
5772
+ src1_type == GGML_TYPE_IQ2_S ||
5773
+ src1_type == GGML_TYPE_IQ3_XXS ||
5774
+ src1_type == GGML_TYPE_IQ1_S ||
5775
+ src1_type == GGML_TYPE_IQ1_M ||
5776
+ src1_type == GGML_TYPE_IQ3_S ||
5777
+ src1_type == GGML_TYPE_IQ4_XS) {
5778
+ return false;
5779
+ }
4792
5780
  }
4793
- if(src0_type == GGML_TYPE_Q4_0 && src1_type == GGML_TYPE_Q4_0) {
4794
- return true;
5781
+
5782
+ if (src1_type == GGML_TYPE_F32) {
5783
+ if (src0_type == GGML_TYPE_Q1_0 ||
5784
+ src0_type == GGML_TYPE_NVFP4 ||
5785
+ src0_type == GGML_TYPE_Q2_K ||
5786
+ src0_type == GGML_TYPE_Q3_K ||
5787
+ src0_type == GGML_TYPE_Q4_K ||
5788
+ src0_type == GGML_TYPE_Q5_K ||
5789
+ src0_type == GGML_TYPE_Q6_K ||
5790
+ src0_type == GGML_TYPE_IQ2_XXS ||
5791
+ src0_type == GGML_TYPE_IQ2_XS ||
5792
+ src0_type == GGML_TYPE_IQ2_S ||
5793
+ src0_type == GGML_TYPE_IQ3_XXS ||
5794
+ src0_type == GGML_TYPE_IQ1_S ||
5795
+ src0_type == GGML_TYPE_IQ1_M ||
5796
+ src0_type == GGML_TYPE_IQ3_S ||
5797
+ src0_type == GGML_TYPE_IQ4_NL ||
5798
+ src0_type == GGML_TYPE_IQ4_XS
5799
+ ) {
5800
+ return false;
5801
+ }
4795
5802
  }
4796
- if(src0_type == GGML_TYPE_Q4_1 && src1_type == GGML_TYPE_Q4_1) {
4797
- return true;
5803
+
5804
+ if (src0_type == src1_type) {
5805
+ if (src1_type == GGML_TYPE_IQ2_XXS ||
5806
+ src1_type == GGML_TYPE_IQ2_XS ||
5807
+ src1_type == GGML_TYPE_IQ2_S ||
5808
+ src1_type == GGML_TYPE_IQ3_XXS ||
5809
+ src1_type == GGML_TYPE_IQ3_S ||
5810
+ src1_type == GGML_TYPE_IQ1_S ||
5811
+ src1_type == GGML_TYPE_IQ1_M) {
5812
+ return false;
5813
+ }
4798
5814
  }
4799
- return false;
5815
+
5816
+ return true;
4800
5817
  }
4801
5818
  case GGML_OP_REPEAT_BACK:
4802
5819
  {
@@ -4828,11 +5845,6 @@ static bool ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const g
4828
5845
  case GGML_OP_COS:
4829
5846
  case GGML_OP_CLAMP:
4830
5847
  case GGML_OP_LOG:
4831
- #if defined (GGML_SYCL_F16)
4832
- return ((op->type == GGML_TYPE_F32 || op->type == GGML_SYCL_F16) && (op->src[0]->type == GGML_TYPE_F32 || op->src[0]->type == GGML_SYCL_F16) && (op->type == op->src[0]->type));
4833
- #else
4834
- return (op->type == GGML_TYPE_F32 && op->src[0]->type == GGML_TYPE_F32) && (op->type == op->src[0]->type);
4835
- #endif
4836
5848
  case GGML_OP_NORM:
4837
5849
  case GGML_OP_L2_NORM:
4838
5850
  case GGML_OP_GROUP_NORM:
@@ -4843,7 +5855,7 @@ static bool ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const g
4843
5855
  case GGML_OP_SCALE:
4844
5856
  return true;
4845
5857
  case GGML_OP_CONT:
4846
- return op->src[0]->type != GGML_TYPE_BF16;
5858
+ return true;
4847
5859
  case GGML_OP_TRI:
4848
5860
  {
4849
5861
  const ggml_tensor * src0 = op->src[0];
@@ -4863,16 +5875,29 @@ static bool ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const g
4863
5875
  case GGML_OP_ROPE:
4864
5876
  case GGML_OP_ROPE_BACK:
4865
5877
  case GGML_OP_IM2COL:
4866
- return true;
5878
+ case GGML_OP_IM2COL_3D:
4867
5879
  case GGML_OP_UPSCALE:
4868
- return op->src[0]->type == GGML_TYPE_F32 && op->op_params[0] == GGML_SCALE_MODE_NEAREST && !(op->op_params[0] & GGML_SCALE_FLAG_ANTIALIAS);
5880
+ return true;
5881
+ case GGML_OP_COL2IM_1D:
5882
+ return ggml_is_contiguous(op->src[0]) &&
5883
+ (op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_F16
5884
+ #ifdef GGML_SYCL_HAS_BF16
5885
+ || op->type == GGML_TYPE_BF16
5886
+ #endif
5887
+ ) &&
5888
+ op->src[0]->type == op->type;
5889
+ case GGML_OP_CONV_3D:
5890
+ return op->type == GGML_TYPE_F32 &&
5891
+ (op->src[0]->type == GGML_TYPE_F32 || op->src[0]->type == GGML_TYPE_F16) &&
5892
+ op->src[1]->type == GGML_TYPE_F32 &&
5893
+ ggml_is_contiguous(op->src[0]) &&
5894
+ ggml_is_contiguous(op->src[1]);
4869
5895
  case GGML_OP_SUM:
4870
5896
  case GGML_OP_SUM_ROWS:
4871
5897
  case GGML_OP_MEAN:
4872
5898
  return ggml_is_contiguous(op->src[0]);
4873
5899
  case GGML_OP_ARGSORT:
4874
- return op->src[0]->ne[0] * sizeof(int) <=
4875
- ggml_sycl_info().devices[device].smpbo;
5900
+ return true;
4876
5901
  case GGML_OP_TOP_K: {
4877
5902
  const ggml_tensor * src0 = op->src[0];
4878
5903
  const int k = op->ne[0];
@@ -4883,15 +5908,14 @@ static bool ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const g
4883
5908
  k > 0 && k <= 32;
4884
5909
  }
4885
5910
  case GGML_OP_POOL_2D:
4886
- return true;
5911
+ case GGML_OP_POOL_1D:
4887
5912
  case GGML_OP_ACC:
4888
- return ggml_is_contiguous(op->src[0]) && ggml_is_contiguous(op->src[1]);
5913
+ return true;
4889
5914
  case GGML_OP_PAD:
4890
- // TODO: add circular padding support for syscl, see https://github.com/ggml-org/llama.cpp/pull/16985
4891
5915
  if (ggml_get_op_params_i32(op, 8) != 0) {
4892
5916
  return false;
4893
5917
  }
4894
- return ggml_is_contiguous(op->src[0]);
5918
+ return true;
4895
5919
  case GGML_OP_LEAKY_RELU:
4896
5920
  case GGML_OP_TIMESTEP_EMBEDDING:
4897
5921
  case GGML_OP_RWKV_WKV6:
@@ -4907,6 +5931,23 @@ static bool ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const g
4907
5931
  return op->type == GGML_TYPE_F32;
4908
5932
  case GGML_OP_ARANGE:
4909
5933
  return op->type == GGML_TYPE_F32;
5934
+ case GGML_OP_SSM_SCAN:
5935
+ if (op->src[3]->ne[0] == 1) {
5936
+ // Mamba2
5937
+ // (kernel only supports (d_state == 128 || d_state == 256) && d_head % WARP_SIZE == 0)
5938
+ return (op->src[0]->ne[0] == 128 || op->src[0]->ne[0] == 256) && op->src[0]->ne[1] % WARP_SIZE == 0;
5939
+ } else {
5940
+ // TODO Mamba-1 not yet ported to SYCL
5941
+ return false;
5942
+ }
5943
+ case GGML_OP_FILL:
5944
+ case GGML_OP_CUMSUM:
5945
+ case GGML_OP_DIAG:
5946
+ case GGML_OP_CROSS_ENTROPY_LOSS:
5947
+ case GGML_OP_CROSS_ENTROPY_LOSS_BACK:
5948
+ return true;
5949
+ case GGML_OP_SOLVE_TRI:
5950
+ return op->src[0]->ne[0] <= SYCL_SOLVE_TRI_MAX_N && op->src[1]->ne[0] <= SYCL_SOLVE_TRI_MAX_K;
4910
5951
  case GGML_OP_FLASH_ATTN_EXT:
4911
5952
  return ggml_sycl_flash_attn_ext_supported(device, op);
4912
5953
  default:
@@ -4916,6 +5957,13 @@ static bool ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const g
4916
5957
  GGML_UNUSED(dev);
4917
5958
  }
4918
5959
 
5960
+ static bool ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, const ggml_tensor * op) {
5961
+ bool res = do_ggml_backend_sycl_device_supports_op(dev, op);
5962
+ GGML_SYCL_DEBUG("[SYCL] call %s op->op=%s op->type=%s -> %s\n", __func__, ggml_op_name(op->op),
5963
+ ggml_type_name(op->type), res ? "true" : "false");
5964
+ return res;
5965
+ }
5966
+
4919
5967
  static bool ggml_backend_sycl_device_supports_buft(ggml_backend_dev_t dev, ggml_backend_buffer_type_t buft) {
4920
5968
  if (buft->iface.get_name != ggml_backend_sycl_buffer_type_get_name) {
4921
5969
  return false;
@@ -5031,6 +6079,250 @@ static ggml_backend_dev_t ggml_backend_sycl_reg_get_device(ggml_backend_reg_t re
5031
6079
  return ctx->devices[index];
5032
6080
  }
5033
6081
 
6082
+ // ==========================================================================
6083
+ // Tensor parallelism (--split-mode tensor) for the SYCL backend.
6084
+ //
6085
+ // The meta-backend invokes these three entry points via get_proc_address:
6086
+ // * ggml_backend_sycl_comm_init - one-time per-graph setup
6087
+ // * ggml_backend_sycl_comm_allreduce_tensor - per-allreduce step
6088
+ // * ggml_backend_sycl_comm_free - tear-down
6089
+ //
6090
+ // For N=2 (dual-GPU), this is a degenerate ring allreduce with dual paths
6091
+ // chosen by tensor size:
6092
+ //
6093
+ // * Small (nelem < 32K): FP32 direct memcpy + per-device ADD
6094
+ // kernel. The kernel depends_on() its corresponding memcpy event
6095
+ // so it doesn't read partial data. Both devices run in parallel.
6096
+ //
6097
+ // * Large (nelem >= 32K): BF16-compressed. Each device compresses
6098
+ // its FP32 partial to BF16 locally, cross-device memcpys
6099
+ // to the peer (half the PCI bandwidth), where it is decompressed
6100
+ // and added into the local FP32 partial. 6 SYCL submissions per
6101
+ // allreduce (2 compress + 2 memcpy + 2 decompress-add) vs the
6102
+ // 4 for the small path, but the bandwidth saving > 6 GB/s PCIe x 2
6103
+ // dominates for larger tensors.
6104
+ //
6105
+ // Storage: A persistent uint8_t buffer per device, sized to
6106
+ // 4 * nelem bytes. Both paths reinterpret the same bytes (small path
6107
+ // as nelem floats; large path as outbox + inbox = 2*nelem uint16_t
6108
+ // each, using the full 4*nelem byte budget either way). Single
6109
+ // alloc+free per device keeps the SYCL pool's strict-LIFO invariant
6110
+ // trivial.
6111
+ //
6112
+ // For non-(N=2 FP32 contiguous) cases, comm_init or comm_allreduce_tensor
6113
+ // returns null/false, causing the meta-backend to use its generic
6114
+ // butterfly all-reduce fallback.
6115
+ // ==========================================================================
6116
+
6117
+ struct ggml_backend_sycl_comm_context {
6118
+ std::vector<ggml_backend_t> backends;
6119
+ // ONE persistent per-device byte buffer, 4*nelem bytes. Both the
6120
+ // FP32 small-tensor path and the BF16 large-tensor path share it
6121
+ // by reinterpreting.
6122
+ std::unique_ptr<ggml_sycl_pool_alloc<uint8_t>> buf0;
6123
+ std::unique_ptr<ggml_sycl_pool_alloc<uint8_t>> buf1;
6124
+ int64_t buf_nelem = 0;
6125
+ };
6126
+
6127
+ void * ggml_backend_sycl_comm_init(ggml_backend_t * backends, size_t n_backends) try {
6128
+ for (size_t i = 0; i < n_backends; ++i) {
6129
+ if (!ggml_backend_is_sycl(backends[i])) {
6130
+ return nullptr;
6131
+ }
6132
+ }
6133
+
6134
+ // Initial version: N=2 only. For N!=2, returning null makes the
6135
+ // meta-backend skip this backend-specific allreduce entirely.
6136
+ if (n_backends != 2) {
6137
+ return nullptr;
6138
+ }
6139
+
6140
+ auto * ctx = new ggml_backend_sycl_comm_context;
6141
+ ctx->backends.assign(backends, backends + n_backends);
6142
+ auto * sctx0 = (ggml_backend_sycl_context *) backends[0]->context;
6143
+ auto * sctx1 = (ggml_backend_sycl_context *) backends[1]->context;
6144
+ ctx->buf0 = std::make_unique<ggml_sycl_pool_alloc<uint8_t>>(sctx0->pool());
6145
+ ctx->buf1 = std::make_unique<ggml_sycl_pool_alloc<uint8_t>>(sctx1->pool());
6146
+ return ctx;
6147
+ }
6148
+ catch (const sycl::exception &) { return nullptr; }
6149
+ catch (...) { return nullptr; }
6150
+
6151
+ void ggml_backend_sycl_comm_free(void * comm_ctx_v) {
6152
+ auto * comm_ctx = static_cast<ggml_backend_sycl_comm_context *>(comm_ctx_v);
6153
+ if (comm_ctx == nullptr) {
6154
+ return;
6155
+ }
6156
+
6157
+ // Sync both per-device queues so the pool_alloc destructors don't
6158
+ // return memory still in use by the last kernel.
6159
+ if (comm_ctx->backends.size() == 2) {
6160
+ auto * sctx0 = (ggml_backend_sycl_context *) comm_ctx->backends[0]->context;
6161
+ auto * sctx1 = (ggml_backend_sycl_context *) comm_ctx->backends[1]->context;
6162
+ try {
6163
+ sctx0->stream()->wait();
6164
+ sctx1->stream()->wait();
6165
+ } catch (...) { /* best effort during shutdown */ }
6166
+ }
6167
+
6168
+ delete comm_ctx;
6169
+ }
6170
+
6171
+ bool ggml_backend_sycl_comm_allreduce_tensor(void * comm_ctx_v, struct ggml_tensor ** tensors) try {
6172
+ if (comm_ctx_v == nullptr) {
6173
+ return false;
6174
+ }
6175
+
6176
+ auto * comm_ctx = static_cast<ggml_backend_sycl_comm_context *>(comm_ctx_v);
6177
+ const size_t n_backends = comm_ctx->backends.size();
6178
+
6179
+ // Fast path: N=2, F32/F16, contiguous, matching shapes.
6180
+ if (n_backends != 2) {
6181
+ return false;
6182
+ }
6183
+ // Accept F32 or F16 inputs natively (types must match). F16 takes the
6184
+ // direct 2-byte memcpy + add path below; other types return false so the
6185
+ // meta-backend uses its generic all-reduce.
6186
+ if (tensors[0]->type != tensors[1]->type) {
6187
+ return false;
6188
+ }
6189
+ if (tensors[0]->type != GGML_TYPE_F32 && tensors[0]->type != GGML_TYPE_F16) {
6190
+ return false;
6191
+ }
6192
+ if (!ggml_is_contiguous(tensors[0]) || !ggml_is_contiguous(tensors[1])) {
6193
+ return false;
6194
+ }
6195
+ if (ggml_nelements(tensors[0]) != ggml_nelements(tensors[1])) {
6196
+ return false;
6197
+ }
6198
+
6199
+ const int64_t nelem = ggml_nelements(tensors[0]);
6200
+ const size_t nbytes = ggml_nbytes(tensors[0]);
6201
+ if (nelem == 0) {
6202
+ return true;
6203
+ }
6204
+
6205
+ auto * ctx0 = (ggml_backend_sycl_context *) comm_ctx->backends[0]->context;
6206
+ auto * ctx1 = (ggml_backend_sycl_context *) comm_ctx->backends[1]->context;
6207
+ queue_ptr q0 = ctx0->stream();
6208
+ queue_ptr q1 = ctx1->stream();
6209
+
6210
+ // Grow per-device byte buffers if needed (4 * nelem bytes each).
6211
+ if (comm_ctx->buf_nelem < nelem) {
6212
+ comm_ctx->buf0->realloc(nelem * 4);
6213
+ comm_ctx->buf1->realloc(nelem * 4);
6214
+ comm_ctx->buf_nelem = nelem;
6215
+ }
6216
+ uint8_t * buf0 = comm_ctx->buf0->get();
6217
+ uint8_t * buf1 = comm_ctx->buf1->get();
6218
+
6219
+ // F16 native path: direct 2-byte cross-device copy + add, skipping the
6220
+ // F32 round-trip the meta-backend fallback would force. Cross-device copies
6221
+ // go through dev2dev_memcpy because the two devices are in separate SYCL
6222
+ // contexts (a raw peer-USM q->memcpy would be a silent no-op).
6223
+ if (tensors[0]->type == GGML_TYPE_F16) {
6224
+ sycl::half * f16_out0 = (sycl::half *) tensors[0]->data;
6225
+ sycl::half * f16_out1 = (sycl::half *) tensors[1]->data;
6226
+ sycl::half * f16_tmp0 = (sycl::half *) buf0;
6227
+ sycl::half * f16_tmp1 = (sycl::half *) buf1;
6228
+
6229
+ q0->wait();
6230
+ q1->wait();
6231
+ dev2dev_memcpy(ctx0->device, *q0, ctx1->device, *q1, f16_tmp0, tensors[1]->data, nbytes);
6232
+ dev2dev_memcpy(ctx1->device, *q1, ctx0->device, *q0, f16_tmp1, tensors[0]->data, nbytes);
6233
+
6234
+ q0->submit([&](sycl::handler & h) {
6235
+ h.parallel_for(sycl::range<1>(nelem), [=](sycl::id<1> i) {
6236
+ f16_out0[i] = (sycl::half) ((float) f16_out0[i] + (float) f16_tmp0[i]);
6237
+ });
6238
+ });
6239
+ q1->submit([&](sycl::handler & h) {
6240
+ h.parallel_for(sycl::range<1>(nelem), [=](sycl::id<1> i) {
6241
+ f16_out1[i] = (sycl::half) ((float) f16_out1[i] + (float) f16_tmp1[i]);
6242
+ });
6243
+ });
6244
+ return true;
6245
+ }
6246
+
6247
+ float * out0 = (float *) tensors[0]->data;
6248
+ float * out1 = (float *) tensors[1]->data;
6249
+
6250
+ // BF16 threshold: above this, the PCIe savings from halving the
6251
+ // cross-device bytes outweigh the 2 extra compress kernels.
6252
+ // Below: stay on the FP32 fast path. Threshold mirrors the CUDA
6253
+ // NCCL allreduce pattern for n_backends=2.
6254
+ static constexpr int64_t BF16_THRESHOLD = 32768;
6255
+
6256
+ if (nelem < BF16_THRESHOLD) {
6257
+ // FP32 small path: 4 SYCL submissions per allreduce.
6258
+ float * tmp0 = (float *) buf0;
6259
+ float * tmp1 = (float *) buf1;
6260
+
6261
+ // COMM-D2D-FIX: the two devices are in SEPARATE SYCL contexts, so a raw
6262
+ // q->memcpy of a peer USM pointer is a silent no-op. Route cross-device
6263
+ // copies through dev2dev_memcpy (L0 direct copy / host staging). It is
6264
+ // synchronous, so wait for the local partials to be produced first.
6265
+ q0->wait();
6266
+ q1->wait();
6267
+ dev2dev_memcpy(ctx0->device, *q0, ctx1->device, *q1, tmp0, tensors[1]->data, nbytes);
6268
+ dev2dev_memcpy(ctx1->device, *q1, ctx0->device, *q0, tmp1, tensors[0]->data, nbytes);
6269
+
6270
+ q0->submit([&](sycl::handler & h) {
6271
+ h.parallel_for(sycl::range<1>(nelem), [=](sycl::id<1> i) {
6272
+ out0[i] += tmp0[i];
6273
+ });
6274
+ });
6275
+ q1->submit([&](sycl::handler & h) {
6276
+ h.parallel_for(sycl::range<1>(nelem), [=](sycl::id<1> i) {
6277
+ out1[i] += tmp1[i];
6278
+ });
6279
+ });
6280
+ return true;
6281
+ }
6282
+
6283
+ // BF16 large path: 6 SYCL submissions per allreduce, but the
6284
+ // cross-device memcpy is HALF the bytes. Pure bit-shift
6285
+ // conversion (no rounding) — matches ggml's truncating fp32->bf16.
6286
+ uint16_t * outbox0 = (uint16_t *) buf0;
6287
+ uint16_t * inbox0 = outbox0 + nelem;
6288
+ uint16_t * outbox1 = (uint16_t *) buf1;
6289
+ uint16_t * inbox1 = outbox1 + nelem;
6290
+
6291
+ // Phase A: compress each device's local partial in parallel.
6292
+ sycl::event c0 = q0->parallel_for(sycl::range<1>(nelem), [=](sycl::id<1> i) {
6293
+ outbox0[i] = (uint16_t) (sycl::bit_cast<uint32_t>(out0[i]) >> 16);
6294
+ });
6295
+
6296
+ sycl::event c1 = q1->parallel_for(sycl::range<1>(nelem), [=](sycl::id<1> i) {
6297
+ outbox1[i] = (uint16_t) (sycl::bit_cast<uint32_t>(out1[i]) >> 16);
6298
+ });
6299
+
6300
+ // Phase B: COMM-D2D-FIX-BF16 cross-device copy of compressed bytes via
6301
+ // dev2dev_memcpy (separate SYCL contexts; sync copy after compress).
6302
+ const size_t bf16_bytes = nelem * sizeof(uint16_t);
6303
+ c0.wait();
6304
+ c1.wait();
6305
+ dev2dev_memcpy(ctx0->device, *q0, ctx1->device, *q1, inbox0, outbox1, bf16_bytes);
6306
+ dev2dev_memcpy(ctx1->device, *q1, ctx0->device, *q0, inbox1, outbox0, bf16_bytes);
6307
+
6308
+ // Phase C: decompress + add into local FP32 partial.
6309
+ q0->submit([&](sycl::handler & h) {
6310
+ h.parallel_for(sycl::range<1>(nelem), [=](sycl::id<1> i) {
6311
+ out0[i] += sycl::bit_cast<float>(((uint32_t) inbox0[i]) << 16);
6312
+ });
6313
+ });
6314
+
6315
+ q1->submit([&](sycl::handler & h) {
6316
+ h.parallel_for(sycl::range<1>(nelem), [=](sycl::id<1> i) {
6317
+ out1[i] += sycl::bit_cast<float>(((uint32_t) inbox1[i]) << 16);
6318
+ });
6319
+ });
6320
+
6321
+ return true;
6322
+ }
6323
+ catch (const sycl::exception &) { return false; }
6324
+ catch (...) { return false; }
6325
+
5034
6326
  static void *ggml_backend_sycl_reg_get_proc_address(ggml_backend_reg_t reg, const char *name) {
5035
6327
  GGML_UNUSED(reg);
5036
6328
 
@@ -5038,6 +6330,17 @@ static void *ggml_backend_sycl_reg_get_proc_address(ggml_backend_reg_t reg, cons
5038
6330
  return (void *)ggml_backend_sycl_split_buffer_type;
5039
6331
  }
5040
6332
 
6333
+ // Tensor parallelism (--split-mode tensor) entry points.
6334
+ if (strcmp(name, "ggml_backend_comm_init") == 0) {
6335
+ return (void *)ggml_backend_sycl_comm_init;
6336
+ }
6337
+ if (strcmp(name, "ggml_backend_comm_free") == 0) {
6338
+ return (void *)ggml_backend_sycl_comm_free;
6339
+ }
6340
+ if (strcmp(name, "ggml_backend_comm_allreduce_tensor") == 0) {
6341
+ return (void *)ggml_backend_sycl_comm_allreduce_tensor;
6342
+ }
6343
+
5041
6344
  // SYCL doesn't support registering host memory, left here for reference
5042
6345
  // "ggml_backend_register_host_buffer"
5043
6346
  // "ggml_backend_unregister_host_buffer"