Compare commits

...
43 Commits
Author SHA1 Message Date
R0CKSTAR 86a2835320 musa: drop mp_21 from the default architectures (#30203)
* musa: drop mp_21 from the default architectures

Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>

* Address review comments

Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>

* Address review comments

Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>

---------

Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
2026-10-09 12:51:24 +02:00
Pascal e94acad853 ci: fix Models Backend webgpu by supporting GGML_OP_DUP (#30216)
The mixed batch path of PR-29622 writes the token rows with set_rows
into a dup of the embeddings. WebGPU did not support DUP, so the dup
ran on the CPU while the set_rows writing into it was scheduled on
WebGPU, which then bound a CPU buffer and crashed. DUP is the same copy
as CPY and CONT and now goes through the same path.
2026-10-09 12:27:32 +02:00
Sigbjørn Skjæret 5e1d74043e ci : set run-name for publish release (#30214) 2026-10-09 11:12:24 +02:00
Xuan-Son Nguyen b42b7e6d30 server: use port 9931 by default (#30159)
* server: use port 9931 by default

* revert unrelated changes
2026-10-09 10:40:03 +02:00
Aleksander Grygier b013e56a71 ui : add the models manager (#29583)
* ui : add the drawer, sheet and grouped list primitives

Add the drawer and sheet overlay components, the toggle and toggle
group, the shared searchable input, the collapsible sections and the
grouped list, and the near viewport helper they position with.

Assisted-by: pi:llama.cpp/DeepSeek-V4.1-Flash

* ui : add the shared model components and data layer

Add the model row pieces shared by the selector and the manager, the
models store with its download status feed, the huggingface and
migration services, and the model utils, enums and constants they
read through.

Assisted-by: pi:llama.cpp/DeepSeek-V4.1-Flash

* ui : rework the models selector around its providers

Group the selector components under their own folder, rework the list
around the provider grouping, add the mobile trigger and the download
item, and derive the selection state in one hook.

Assisted-by: pi:llama.cpp/DeepSeek-V4.1-Flash

* ui : add the models manager

Add the manager table with its repo, quant and status rows, the
filters and the toolbar, the row actions in a drawer, the downloads
section, the manage models dialog and the stories and tests.

Assisted-by: pi:llama.cpp/DeepSeek-V4.1-Flash

* ui : rework the chat form actions and the mobile experience

Move the add actions into a drawer and a sheet, float the model
actions in a bar on a phone, open the context panel in a drawer, and
derive the attachment, reasoning and tools menus in hooks.

Assisted-by: pi:llama.cpp/DeepSeek-V4.1-Flash

* ui : move the mcp servers to the sidebar rail and tidy the dialogs

Move the mcp servers dialog to the sidebar rail with its menu
entries, even out the dialogs on a phone, and group the mcp and
settings stores behind their own modules.

Assisted-by: pi:llama.cpp/DeepSeek-V4.1-Flash

* ui : polish the shell

The secondary button gets its own look back and the chat add button its
own light surface. Pressable elements get a pointer cursor again, the
font rendering smooths on the app shell, and the agent skills stay out
of prettier's way.

The root layout props probe rework that came with this polish reads the
providers api url and stays with the providers change instead.

Assisted-by: pi:llama.cpp/DeepSeek-V4.1-Flash

* ui : wrap the model id classes

Assisted-by: pi:llama.cpp/DeepSeek-V4.1-Flash

* ui : load the model the new chat CTA picks

Start a new chat selected the model but left it unloaded, so the chat
opened against a server with nothing in memory. Both CTAs share a load
helper now.

Assisted-by: pi:llama.cpp/DeepSeek-V4.1-Flash

* refactor: Post-review fixes
2026-10-09 10:34:45 +02:00
Pascal 2c564f43df ci: fix HIP quality check by ignoring the 320/256 FA kernel spill (#30212) 2026-10-09 10:29:22 +02:00
Pascal 1f8fa52318 metal: add the 128/96 flash attention kernels, accept every compiled pair (#30209) 2026-10-09 10:15:32 +02:00
ynankani 8a1a9b5126 CUDA: pass src1 precision to host MMQ config helpers (#30168)
* CUDA: pass src1 precision to host MMQ config helpers

Signed-off-by: ynankani <ynankani@nvidia.com>

* make src1 prec explicit

Signed-off-by: ynankani <ynankani@nvidia.com>

---------

Signed-off-by: ynankani <ynankani@nvidia.com>
2026-10-09 12:02:52 +05:30
Mark TodorovichandMax Krasnyansky d4d82d67f4 hexagon: fix IM2COL patch-embed DMA ring overflow (#30189)
* hexagon: fix IM2COL patch-embed DMA ring overflow

The exact-tiling (stride == kernel, no pad/dilation) IM2COL DMA kernel
issues IC*KH DDR->VTCM descriptors per output row without checking the
return value of dma_queue_push(), and then pops IC*KH times. The per-thread
DMA ring holds 256 entries and a push into a full ring returns false
and drops the transfer, so for IC*KH > 255 the remaining rows of the
VTCM staging buffer were never written and stale data (often NaN/inf)
leaked into the output.

Solution is to retire the oldest descriptor when the ring is full, just as the blocked
kernel in the same file already does, and wait with dma_queue_flush().

For testing, added exact-tiling test cases with IC*KH > 256 (2D 1x1, 2D 2x2 patch
embed and 1D, F16 and F32 dst), which fail on HTP without this fix.

* Apply suggestion from @max-krasnyansky

---------

Co-authored-by: Max Krasnyansky <maxk@qti.qualcomm.com>
2026-10-09 08:04:09 +02:00
Nicolas Mowen 3d65c90d04 sycl : Q5_K reorder-layout MMVQ and fused GLU (#29375) 2026-10-08 22:19:32 -04:00
bri-prism de7fa0a3c6 Musa FWHT fix (#30167) 2026-10-08 21:56:54 +02:00
71ad0590f4 CUDA: improve top-k algorithm selection (#28713)
* CUDA: radix top-k for large row counts

Replaces CUB's per-row DeviceTopKKernel with a grid-over-rows radix select,
gated on GGML_CUDA_TOPK_RADIX_MIN_ROWS. On qwen4exp at 34,816 tokens this cuts
top-k from 1,671,253 launches / 5,761.8 ms to 2,329 / 941.8 ms.

* CUDA: select the TOP_K implementation by shape

Replace the nrows/ncols special case with the decision boundary from #28547
(as implemented in #29278): bitonic for short rows, radix select for several
long rows, and DeviceTopK or CUB argsort for a single long row. The
thresholds stay overridable at build time.

Two refinements on top of that boundary:
- bitonic stays in use for rows up to a padded 1024 while the rows fit in one
  wave of blocks (nrows <= number of SMs); radix select pays a fixed cost of
  about a dozen launches that only amortizes over more rows
- with DeviceTopK available, it handles up to two rows

Radix select now processes rows in chunks so its scratch memory stays bounded,
and the bitonic path keeps its chunking. HIP and MUSA keep their previous
thresholds.

Add perf cases around the bitonic/radix crossover to test-backend-ops.

* CUDA: make top-k comments less verbose

* CUDA: remove the TOP_K width limit from supports_op

* CUDA: use DeviceTopK for single-row TOP_K if available

* CUDA: avoid ncols overflow in the TOP_K bitonic check

* CUDA: share the row chunking helper between argsort and top-k

* CUDA: do the TOP_K radix blocks_per_row math in int64_t

* CUDA: rename GGML_CUDA_TOP_K_NROWS_THRESHOLD_DEVICETOPK to GGML_CUDA_TOP_K_NROWS_THRESHOLD

* CUDA: share one sort helper between the bitonic and CUB TOP_K paths

* CUDA: update the TOP_K TODO, threshold and chunking comments

* tests: add TOP_K cases that span several row chunks

* CUDA: use int64_t col in the TOP_K radix loops, fix threshold comment

* CUDA: limit TOP_K and ARGSORT support to ne[0] <= INT_MAX

---------

Co-authored-by: praneshgo <227579474+praneshgo@users.noreply.github.com>
Co-authored-by: Pranesh Gonegandla <pgonegandla@nvidia.com>
2026-10-08 18:34:27 +02:00
Hrishith Thadicherla a11f57ba93 model : fix DFlash output head sharing (#30111)
* llama : fix DFlash output head sharing

Assisted-by: Codex

* dflash : read tied output weights from GGUF metadata

Assisted-by: Codex

* llama : share tied word embedding metadata

Assisted-by: Codex

* llama : remove DFlash embedding head fallback

Assisted-by: Codex
2026-10-08 17:42:28 +02:00
Johannes Gäßler fc9ce6b9d5 CUDA: fix MMQ out-of-bounds reads (#29953) 2026-10-08 16:41:49 +02:00
Prabhsimran Singh c35b66744f CUDA : looped PAD kernel for more than 65535 rows or slices (#30147) 2026-10-08 15:36:53 +02:00
hey-gmandOliver Simons 1167d3f42c CUDA: fix CCCL version guard breaking on major version rollover (#29453)
* CUDA: fix CCCL version guard breaking on major version rollover

The guard compared the major and minor components independently:

    CCCL_MAJOR_VERSION >= 3 && CCCL_MINOR_VERSION >= 1

Minor resets to 0 whenever a new major series is cut, so on CCCL 4.x
this evaluates as 4 >= 3 && 0 >= 1, i.e. false. STRIDED_ITERATOR_AVAILABLE
stops being defined and argsort silently falls back to the
init_offsets path. Nothing warns and the build still succeeds, so the
regression is a quiet performance loss rather than a compile error.

CCCL already exposes the version as a single packed integer in
MMMmmmpp form, which is what its own version header uses:

    CCCL_VERSION = MAJOR * 1000000 + MINOR * 1000 + PATCH

so 3.4.3 is 3004003 and ">= 3.1" is a plain ">= 3001000". One
comparison, with no component arithmetic left to get wrong.

Checked against a hand-written "version >= 3.1" reference over 2.9.9,
3.0.0, 3.1.0, 3.1.99, 3.2.0, 3.4.3, 3.9.9, 3.99.99, 4.0.0, 4.2.7 and
5.0.0: no divergences. The old guard disagreed at 4.0.0 and 5.0.0.

Verified on RTX 4070 (sm_89), CUDA 13.4, CCCL 3.4.3:

  - cmake --build build --config Release: exit 0
  - test-backend-ops test -o ARGSORT -b CUDA0: 98/98 passed, CUDA0 OK

Note that a passing regression test does not on its own prove the guard
is still taken, since the fallback path passes too. Preprocessing the
real translation unit confirms the strided-iterator branch is the one
compiled in: counting_iterator is present, init_offsets is not.

Signed-off-by: Heitor <heitorgm@outlook.com>

* Update ggml/src/ggml-cuda/argsort.cu

* Apply suggestion from @ORippler

---------

Signed-off-by: Heitor <heitorgm@outlook.com>
Co-authored-by: Oliver Simons <osimons@nvidia.com>
2026-10-08 15:29:34 +02:00
Simon Sudarushkin 4f92965a7b ui: apply ui_settings on first visit in router mode (#29668) 2026-10-08 15:08:22 +02:00
Aman Gupta c811cb8f0a llama: support MoE cache over multiple GPUs (#30112) 2026-10-08 15:50:49 +05:30
Leebr Data ConsultingandIgor Okulist 033df86b69 server : preserve context checkpoints across slot save/restore (#26004)
* server : preserve context checkpoints across slot save/restore

Append the checkpoints after the packed server_tokens payload added in #26640
and count them in n_written / n_read, so a restored slot can still roll back to
a checkpoint instead of re-processing the whole prompt.

* server : drop draft checkpoint data that does not match the draft context

Restoring a slot saved with a different draft KV cache type aborted in
load_dft(). Test-load one draft checkpoint on restore and drop the draft
data if it does not fit, instead of crashing. Adds a regression test.

Co-authored-by: Igor Okulist <okigan@gmail.com>

* server : harden the checkpoint appendix of slot save files

Bound each blob size by the bytes left in the file before allocating, open the
file with UTF-8 paths on Windows like the llama state payload, fall back to full
prompt re-processing when a checkpoint restored from a slot file fails to load,
and replace the 1024 count cap by keeping the last n_ctx_checkpoints while reading.

* server : report an incomplete checkpoint appendix as a failed slot save

Return an error to the client when the appendix cannot be written, like a
failed payload write, and make the oversized-blob test declare a size that
cannot be allocated, so an unbounded allocation fails the test.

* server : reject an empty target state in the checkpoint appendix

A saved checkpoint always holds a target state, an empty blob would roll back
without restoring anything. Also log with the slot id, and load the draft test
model from the HF cache instead of a second download.

* common : return bool from checkpoint load_tgt / load_dft

A checkpoint restored from a slot file falls back to full prompt re-processing
when it fails to load, a checkpoint created in memory still aborts.

---------

Co-authored-by: Igor Okulist <okigan@gmail.com>
2026-10-08 13:18:20 +03:00
gianni-cor ff5888f999 vulkan : fix TOP_K for +inf/NaN inputs and k = 1 on negative values (#30107)
The bucket search in topk_nary_search.comp started from the range
[0, 0xFF800000), which ends just below the ordered-uint mapping of +inf,
so +inf and NaN were never counted. A workgroup block with fewer than k
countable values left the ballot empty and the shader read uninitialized
shared state (hang/device lost on NVIDIA, wrong indices on AMD), and a few
+inf in a block were selected without being counted, dropping real top
values.

Map NaN to -inf on input, start from [0, 0xFFFFFFFF) so every value is
counted, and clamp the top bucket's end (2^32) instead of wrapping to 0.

The k = 1 path compared float bits as signed integers, which orders
negative values backwards; compare floats instead.

Add test_top_k_inf to test-backend-ops: negative values, fewer than k
+inf and many -inf, for k = 1, 10, 40.

Assisted-by: Claude Opus 5.5
2026-10-08 12:14:35 +02:00
Jeff Bolz dac3087394 vulkan: extend sparse FA support to coopmat2 (#30003) 2026-10-08 11:36:20 +02:00
Adrien Gallouët 03aa006acb vendor : update cpp-httplib to 0.60.1 (#30134)
Signed-off-by: Adrien Gallouët <angt@huggingface.co>
2026-10-08 10:57:47 +02:00
Muhammad SaadandJohannes Gäßler 08246a28f6 cuda : support arbitrary striding for unary ops on f16, f32, and bf16 (#29781)
* cuda : support arbitrary striding for unary ops on f16, f32, and bf16

* Remove added newline

---------

Co-authored-by: Johannes Gäßler <johannesg@5d6.de>
2026-10-08 09:21:42 +02:00
bri-prism 46baf1f1fe sycl: FWHT optimizations (#29605) 2026-10-08 02:50:22 -04:00
Max Krasnyansky 097f5b5332 hex-mmadd: do not assume aligned read/write when bias-add is fused (#30133) 2026-10-08 09:48:06 +03:00
Bertay Eren d888016041 cuda : Use byte strides for roll to allow non-contiguous ROLL operations (#29547) 2026-10-08 09:46:38 +03:00
Eden fda1866613 CUDA: fix norm family kernels when ne[2]/ne[3] exceed grid dim limits (#28175) 2026-10-08 09:46:21 +03:00
Titaniumtown ff30363a0e sycl: fuse the delta-net alpha gate (add + unary + mul) (#29687)
* sycl: fuse the delta-net alpha gate (add + unary + mul)

* tests: cover the fused add + unary + mul chain

* sycl: give the fused alpha gate a flat path and pin the node skip
2026-10-08 09:45:21 +03:00
cwriterandcwriter 000bee54a5 sycl: stage bulk uploads (model loading) through a pinned ring buffer (#29608)
Co-authored-by: cwriter <cwriter@localhost>
2026-10-08 02:30:00 -04:00
bosh 37ac634566 model : support classifier_activation for rerankers (#29692)
* model : support classifier_activation for rerankers

Assisted-by: Claude Opus 5.5

* model : map classifier gelu to gelu_erf and accept tanh

Assisted-by: Claude Opus 5.5

* model : default act_cls to tanh, ModernBERT falls back to gelu_erf

Assisted-by: Claude Opus 5.5
2026-10-08 09:29:20 +03:00
SXX 24e41838e0 ggml-cuda: assign four GDN state columns per warp (#30087)
* ggml-cuda: assign two GDN state columns per warp

* ggml-cuda: use 4 GDN state columns per warp at S_v=128

* ggml-cuda: default cols_per_warp=4

* ggml-cuda: address GDN review nits
2026-10-08 11:37:34 +05:30
cwriterandcwriter 847f447c31 sycl: add grouped MoE XMX GEMM (#29245)
Co-authored-by: cwriter <cwriter@localhost>
2026-10-08 01:41:13 -04:00
Sam Malayek 75118a3a59 convert : support Qwen3.5 embedding models (#27920)
* Update convert_hf_to_gguf: support Qwen3.5 embedding models

* Behavior-preserving refactor for conventions.

* conversion: simplify pooling comment
2026-10-08 08:40:01 +03:00
Titaniumtown 9b4ed0ca57 sycl: remove duplicate block-size defines from op headers (#29507)
All were duplicates of the values in ggml/src/ggml-sycl/presets.hpp
2026-10-08 08:39:25 +03:00
Max KrasnyanskyandJhen-Jie Hong 9c2e0e491a hexagon: enable alloc_buffer_n (#30126)
* hex-bufs: add support for alloc_buffer_n

* hex-bufs: add support for splitting large tensors into separate buffers

* hex-bufs: update GGML_HEXAGON_MBUF to accept three values dyn,static,total

* hex-bufs: bump dyn. default to 512MB since 128MB causes perf regressions with big MOEs

* hex-run: add --no-embd-offload option to simplify command lines on devices that need it

* Update scripts/snapdragon/run.py

Co-authored-by: Jhen-Jie Hong <iainst0409@gmail.com>

---------

Co-authored-by: Jhen-Jie Hong <iainst0409@gmail.com>
2026-10-07 18:30:06 -07:00
kurquhar 06cad0b9e7 hexagon: Q6_K weight dequant speedup (#30121)
* hexagon: Q6_K weight dequant speedup

Assisted-by: OpenCode

* unroll by another factor of 2

Assisted-by: OpenCode
2026-10-07 17:13:21 -07:00
Tarek Dakhran a657f7e981 model : add LiquidAI/d1-omni-600M decision model (#30114)
* model : add LiquidAI/d1-omni-600M decision model

Assisted-by: Claude Opus 5.5

* mtmd : keep conformer GLU sigmoid on CUDA

Assisted-by: Claude Opus 5.5

* server : take d1omni audio through images and input_audio, scope memory-less lfm2 to non-causal

Assisted-by: Claude Opus 5.5

* common : rename decision type d1omni to lfm2-d1-omni, server : make images an alias of files

Assisted-by: Claude Opus 5.5
2026-10-08 02:07:49 +02:00
kurquhar aa5e0092fd hexagon: support tiled Q4_K and Q6_K GET_ROWS (#30115)
* hexagon: support tiled Q4_K GET_ROWS

Assisted-by: OpenCode

* properly reject Q4_K views

Assisted-by: OpenCode

* hexagon: support tiled Q6_K GET_ROWS

Assisted-by: OpenCode
2026-10-07 17:04:45 -07:00
kurquhar 70815103c8 hexagon: improved GELU accuracy (#30104)
Assisted-by: OpenCode
2026-10-07 16:57:09 -07:00
Nicolas Budyn bd4eeaa047 chat: fix jinja parser for TranslateGemma (#30096)
* chat: fix jinja parser for translategemma

* log if missing soruce_lang_code or target_lang_code
2026-10-07 22:48:06 +02:00
Frost-54andAlde Rojas 5de733437b chat : name tool and argument parser rules by index (#30088)
* bugfix: infinite recursion caused by a tool named 'call'(#29967)

* chat : use index for schema and argument rules

* tests : remove tests

* tests : add expect_rules to peg test parser

---------

Co-authored-by: Alde Rojas <hello@alde.dev>
2026-10-07 22:37:51 +02:00
Tarek Dakhran 88dcc460d6 model : add LiquidAI/d1-3B decision model (#30110)
* model : add LiquidAI/d1-3b decision model

mtmd : read LFM2 image resize algo from GGUF

Assisted-by: Claude Opus 5.5

* common : rename decision type d1 to lfm2-d1

Assisted-by: Claude Opus 5.5
2026-10-07 22:04:39 +02:00
bri-prism b86d2f0754 cuda: FWHT kernels for block widths above 512 (#29100)
The CUDA FWHT covers widths 64 to 512. It runs one row per warp and keeps N/32
values per lane, so wider blocks need more registers per lane than that layout
allows.

fwht_cuda_block runs one row per thread block with 256 threads, so each thread
keeps N/256 values. Stages below the warp width still shuffle, those up to the
block width go through shared memory, and the rest stay in registers. Same
butterfly and sign convention as the warp kernel.

Widths 64 to 512 keep the warp kernel. 1024 through 8192 use the new one, for
both F32 and F16 sources. ggml_cuda_op_mul_mat_use_fwht (the shared
supports_op/dispatch predicate added in #29096) does not check width, so it
needed no change here: any width it admits that ggml_cuda_op_fwht can't serve
already falls through correctly to the cuBLAS path.

Rebased onto current master with #29096's F16 commit underneath it, since this
depends on the same F16 template infrastructure; that commit applied cleanly,
the only conflict was in test-backend-ops.cpp where an unrelated intervening
commit's own test additions landed near this block.

test-backend-ops on an A10 (lambdalabs): MUL_MAT 1297/1297, including all
FWHT/Hadamard cases (18 existing, 4 new F32 wide, 4 new F16 wide, 2 new
many-rows, 1 too-big boundary moved to 16384).
2026-10-07 20:56:18 +02:00
352 changed files with 14509 additions and 4040 deletions
+1
View File
@@ -155,6 +155,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
FROM base AS server
ENV LLAMA_ARG_HOST=0.0.0.0
ENV LLAMA_ARG_PORT=8080
COPY --from=build /app/full/llama /app/full/llama-server /app
+1
View File
@@ -114,6 +114,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
FROM base AS server
ENV LLAMA_ARG_HOST=0.0.0.0
ENV LLAMA_ARG_PORT=8080
COPY --from=build /app/full/llama /app/full/llama-server /app
+1
View File
@@ -123,6 +123,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
FROM base AS server
ENV LLAMA_ARG_HOST=0.0.0.0
ENV LLAMA_ARG_PORT=8080
COPY --from=build /app/full/llama /app/full/llama-server /app
+1
View File
@@ -151,6 +151,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
FROM base AS server
ENV LLAMA_ARG_HOST=0.0.0.0
ENV LLAMA_ARG_PORT=8080
COPY --from=build /app/lib/ /app
COPY --from=build /app/full/llama /app/full/llama-server /app
+1
View File
@@ -130,6 +130,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
FROM base AS server
ENV LLAMA_ARG_HOST=0.0.0.0
ENV LLAMA_ARG_PORT=8080
COPY --from=build /app/full/llama /app/full/llama-server /app
+1
View File
@@ -227,6 +227,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
FROM base AS server
ENV LLAMA_ARG_HOST=0.0.0.0
ENV LLAMA_ARG_PORT=8080
COPY --from=build /app/full/llama /app/full/llama-server /app/
+1
View File
@@ -136,6 +136,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
FROM base AS server
ENV LLAMA_ARG_HOST=0.0.0.0
ENV LLAMA_ARG_PORT=8080
COPY --from=build /app/full/llama /app/full/llama-server /app
+1
View File
@@ -133,6 +133,7 @@ ENTRYPOINT [ "/llama.cpp/bin/llama-cli" ]
FROM base AS server
ENV LLAMA_ARG_HOST=0.0.0.0
ENV LLAMA_ARG_PORT=8080
WORKDIR /llama.cpp/bin
+1
View File
@@ -117,6 +117,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
FROM base AS server
ENV LLAMA_ARG_HOST=0.0.0.0
ENV LLAMA_ARG_PORT=8080
COPY --from=build /app/full/llama /app/full/llama-server /app
+1
View File
@@ -107,6 +107,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
FROM base AS server
ENV LLAMA_ARG_HOST=0.0.0.0
ENV LLAMA_ARG_PORT=8080
COPY --from=build /app/full/llama /app/full/llama-server /app
+2
View File
@@ -9,6 +9,8 @@ on:
branches:
- master
run-name: "Publish ${{ github.event.workflow_run.display_title }}"
cache-mode: none
permissions:
actions: read
+2 -2
View File
@@ -1479,7 +1479,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
));
add_opt(common_arg(
{"--server-base"}, "URL",
string_format("connect to this server instead of starting a new one, example: 'http://localhost:8080' (default: none)"),
string_format("connect to this server instead of starting a new one, example: 'http://localhost:9931' (default: none)"),
[](common_params & params, const std::string & value) {
params.server_base = value;
}
@@ -2778,7 +2778,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
).set_env("LLAMA_ARG_N_CPU_MOE"));
add_opt(common_arg(
{"--moe-cache-mib"}, "N",
"GPU cache size in MiB for the MoE experts kept in the CPU (default: 0, disabled)",
"GPU cache size in MiB for the MoE experts kept in the CPU. with multiple GPUs, it is split among them like the layers (--tensor-split) (default: 0, disabled)",
[](common_params & params, int value) {
if (value < 0) {
throw std::invalid_argument("invalid value");
+8 -8
View File
@@ -291,7 +291,7 @@ common_peg_parser analyze_tools::build_tool_parser_tag_json(parser_build_context
common_peg_parser tool_choice = p.choice();
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & func = tool.at("function");
std::string name = func.at("name");
const auto schema = common_chat_tool_parameters(func);
@@ -308,7 +308,7 @@ common_peg_parser analyze_tools::build_tool_parser_tag_json(parser_build_context
}
have_call_id = true;
}
auto args_parser = p.tool_args(p.schema(p.json(), "tool-" + name + "-schema", schema));
auto args_parser = p.tool_args(p.schema(p.json(), "tool-" + std::to_string(tool_index) + "-schema", schema));
if (!arguments.start.empty()) {
args_parser = p.literal(arguments.start) + args_parser;
}
@@ -318,7 +318,7 @@ common_peg_parser analyze_tools::build_tool_parser_tag_json(parser_build_context
auto atomic_peek = !arguments.start.empty() ? std::optional(p.peek(p.literal(arguments.start))) : std::nullopt;
auto func_parser = build_func_parser(p, name, call_id_section, have_call_id, args_parser, atomic_peek);
tool_choice |= p.rule("tool-" + name, func_parser);
tool_choice |= p.rule("tool-" + std::to_string(tool_index), func_parser);
});
auto require_calls = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED;
@@ -364,14 +364,14 @@ common_peg_parser analyze_tools::build_tool_parser_tag_tagged(parser_build_conte
common_peg_parser tool_choice = p.choice();
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & func = tool.at("function");
std::string name = func.at("name");
// Build parser for each argument, separating required and optional
std::vector<common_peg_parser> required_parsers;
std::vector<common_peg_parser> optional_parsers;
foreach_parameter(func, [&](const common_chat_schema_property & param, const common_chat_schema_document_ptr & doc) {
foreach_parameter(func, [&](size_t param_index, const common_chat_schema_property & param, const common_chat_schema_document_ptr & doc) {
auto arg =
p.tool_arg(p.tool_arg_open(arguments.name_prefix + p.tool_arg_name(p.literal(param.name)) +
arguments.name_suffix) +
@@ -380,10 +380,10 @@ common_peg_parser analyze_tools::build_tool_parser_tag_tagged(parser_build_conte
p.ac(p.tool_arg_string_value(until_suffix) +
p.tool_arg_close(p.literal(arguments.value_suffix)), arguments.value_suffix) :
(p.tool_arg_json_value(p.schema(
p.json(), "tool-" + name + "-arg-" + param.name + "-schema", doc, *param.schema)) +
p.json(), "tool-" + std::to_string(tool_index) + "-arg-" + std::to_string(param_index) + "-schema", doc, *param.schema)) +
p.tool_arg_close(p.literal(arguments.value_suffix)))));
auto named_arg = p.rule("tool-" + name + "-arg-" + param.name, arg);
auto named_arg = p.rule("tool-" + std::to_string(tool_index) + "-arg-" + std::to_string(param_index), arg);
if (param.required) {
required_parsers.push_back(named_arg);
} else {
@@ -434,7 +434,7 @@ common_peg_parser analyze_tools::build_tool_parser_tag_tagged(parser_build_conte
auto atomic_peek = (!arguments.name_prefix.empty() && !required_parsers.empty()) ?
std::optional(p.peek(p.literal(arguments.name_prefix))) : std::nullopt;
auto func_parser = build_func_parser(p, name, call_id_section, have_call_id, args_seq, atomic_peek);
tool_choice |= p.rule("tool-" + name, func_parser);
tool_choice |= p.rule("tool-" + std::to_string(tool_index), func_parser);
});
auto require_tools = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED;
+20 -14
View File
@@ -483,7 +483,9 @@ common_peg_parser common_chat_peg_builder::standard_constructed_tools(
// Build tool choices for tagged format
auto tool_choices = choice();
for (const auto & tool_def : tools) {
for (size_t i = 0; i < tools.size(); i++) {
const auto & tool_def = tools[i];
if (!tool_def.contains("function")) {
continue;
}
@@ -513,7 +515,7 @@ common_peg_parser common_chat_peg_builder::standard_constructed_tools(
auto tool_parser = tool(tool_open(literal(func_opener) + tool_name(literal(name)) + literal(func_name_suffix)) +
space() + tool_args(args) + space() + tool_close(literal(func_closer)));
tool_choices |= rule("tool-" + name, tool_parser);
tool_choices |= rule("tool-" + std::to_string(i), tool_parser);
}
// Build the section with markers
@@ -560,7 +562,8 @@ common_peg_parser common_chat_peg_builder::python_style_tool_calls(
auto tool_choices = choice();
for (const auto & tool_def : tools) {
for (size_t i = 0; i < tools.size(); i++) {
const auto & tool_def = tools[i];
if (!tool_def.contains("function")) {
continue;
}
@@ -607,7 +610,7 @@ common_peg_parser common_chat_peg_builder::python_style_tool_calls(
space() + tool_args(args) + space() + tool_close(literal(")"))
);
tool_choices |= rule("tool-" + name, tool_parser);
tool_choices |= rule("tool-" + std::to_string(i), tool_parser);
}
if (parallel_tool_calls) {
@@ -635,7 +638,8 @@ common_peg_parser common_chat_peg_builder::build_json_tools_function_is_key(
auto tool_choices = choice();
for (const auto & tool_def : tools) {
for (size_t i = 0; i < tools.size(); i++) {
const auto & tool_def = tools[i];
if (!tool_def.contains("function")) {
continue;
}
@@ -668,10 +672,10 @@ common_peg_parser common_chat_peg_builder::build_json_tools_function_is_key(
// Arguments — either wrapped in args_key or parsed directly
common_peg_parser args_parser = eps();
if (args_key.empty()) {
args_parser = tool_args(schema(json(), "tool-" + name + "-schema", params));
args_parser = tool_args(schema(json(), "tool-" + std::to_string(i) + "-schema", params));
} else {
args_parser = literal("\"" + effective_args_key + "\"") + space() + literal(":") + space() +
tool_args(schema(json(), "tool-" + name + "-schema", params));
tool_args(schema(json(), "tool-" + std::to_string(i) + "-schema", params));
}
inner_fields.push_back(args_parser);
@@ -698,7 +702,7 @@ common_peg_parser common_chat_peg_builder::build_json_tools_function_is_key(
space() + tool_close(literal("}"))
);
tool_choices |= rule("tool-" + name, tool_parser);
tool_choices |= rule("tool-" + std::to_string(i), tool_parser);
}
return tool_choices;
@@ -721,7 +725,8 @@ common_peg_parser common_chat_peg_builder::build_json_tools_nested_keys(
std::string nested_name_field = !name_spec.first.empty() ? name_spec.second : effective_name_key;
std::string nested_args_field = !args_spec.first.empty() ? args_spec.second : effective_args_key;
for (const auto & tool_def : tools) {
for (size_t i = 0; i < tools.size(); i++) {
const auto & tool_def = tools[i];
if (!tool_def.contains("function")) {
continue;
}
@@ -732,7 +737,7 @@ common_peg_parser common_chat_peg_builder::build_json_tools_nested_keys(
auto nested_name = literal("\"" + nested_name_field + "\"") + space() + literal(":") + space() +
atomic(literal("\"") + tool_name(literal(name)) + literal("\""));
auto nested_args = literal("\"" + nested_args_field + "\"") + space() + literal(":") + space() +
tool_args(schema(json(), "tool-" + name + "-schema", params));
tool_args(schema(json(), "tool-" + std::to_string(i) + "-schema", params));
auto nested_object = literal("{") + space() +
nested_name + space() + literal(",") + space() +
@@ -770,7 +775,7 @@ common_peg_parser common_chat_peg_builder::build_json_tools_nested_keys(
auto nested_field = literal("\"" + nested_prefix + "\"") + space() + literal(":") + space() + nested_object;
tool_parser_body = tool_parser_body + nested_field + space() + tool_close(literal("}"));
tool_choices |= rule("tool-" + name, tool(tool_parser_body));
tool_choices |= rule("tool-" + std::to_string(i), tool(tool_parser_body));
}
return tool_choices;
@@ -790,7 +795,8 @@ common_peg_parser common_chat_peg_builder::build_json_tools_flat_keys(
auto name_key_parser = literal("\"" + effective_name_key + "\"");
auto args_key_parser = literal("\"" + effective_args_key + "\"");
for (const auto & tool_def : tools) {
for (size_t i = 0; i < tools.size(); i++) {
const auto & tool_def = tools[i];
if (!tool_def.contains("function")) {
continue;
}
@@ -801,7 +807,7 @@ common_peg_parser common_chat_peg_builder::build_json_tools_flat_keys(
auto tool_name_ = name_key_parser + space() + literal(":") + space() +
atomic(literal("\"") + tool_name(literal(name)) + literal("\""));
auto tool_args_ = args_key_parser + space() + literal(":") + space() +
tool_args(schema(json(), "tool-" + name + "-schema", params));
tool_args(schema(json(), "tool-" + std::to_string(i) + "-schema", params));
// Build ID parsers if keys are provided
common_peg_parser id_parser = eps();
@@ -861,7 +867,7 @@ common_peg_parser common_chat_peg_builder::build_json_tools_flat_keys(
}
ordered_body = ordered_body + space() + tool_close(literal("}"));
tool_choices |= rule("tool-" + name, tool(ordered_body));
tool_choices |= rule("tool-" + std::to_string(i), tool(ordered_body));
}
return tool_choices;
+7
View File
@@ -1223,6 +1223,13 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
return common_chat_params_init_minicpm5(tmpl, params);
}
// TranslateGemma - user content must follow a custom schema with language codes
if (src.find("[source_lang_code]") != std::string::npos &&
src.find("[target_lang_code]") != std::string::npos) {
LOG_DBG("Using specialized template: TranslateGemma\n");
return common_chat_params_init_translate_gemma(tmpl, params);
}
// Qwen3-Coder XML tool calls, also used by Nemotron Nano 3, Qwen3.5 and StepFun-3.5-Flash
if (src.find("<tool_call>") != std::string::npos &&
src.find("<function=") != std::string::npos &&
+12 -13
View File
@@ -1170,6 +1170,8 @@ static const std::map<common_decision_type, std::string> COMMON_DECISION_TYPE_NA
{ COMMON_DECISION_TYPE_LAYA, "laya" },
{ COMMON_DECISION_TYPE_CLEF, "clef" },
{ COMMON_DECISION_TYPE_PPLX_DECIDER, "pplx-decider" },
{ COMMON_DECISION_TYPE_LFM2_D1, "lfm2-d1" },
{ COMMON_DECISION_TYPE_LFM2_D1_OMNI, "lfm2-d1-omni" },
};
static common_decision_type common_decision_type_from_string(const std::string & str) {
@@ -1283,7 +1285,8 @@ common_init_result::common_init_result(common_params & params, bool model_only)
// these decision models return a score for each token via the embeddings output
// TODO: maybe improve this in the future
const auto decision_type = common_get_decision_type(model);
if (decision_type == COMMON_DECISION_TYPE_LAYA || decision_type == COMMON_DECISION_TYPE_KEV || decision_type == COMMON_DECISION_TYPE_CLEF) {
if (decision_type == COMMON_DECISION_TYPE_LAYA || decision_type == COMMON_DECISION_TYPE_KEV || decision_type == COMMON_DECISION_TYPE_CLEF ||
decision_type == COMMON_DECISION_TYPE_LFM2_D1_OMNI) {
params.embedding = true;
params.pooling_type = LLAMA_POOLING_TYPE_NONE;
@@ -2385,40 +2388,36 @@ void common_prompt_checkpoint::update_dft(
}
}
void common_prompt_checkpoint::load_tgt(
bool common_prompt_checkpoint::load_tgt(
llama_context * ctx,
llama_seq_id seq_id,
llama_state_seq_flags flags) const {
if (ctx == nullptr) {
return;
return true;
}
if (data_tgt.empty()) {
return;
return true;
}
const size_t n = llama_state_seq_set_data_ext(ctx, data_tgt.data(), data_tgt.size(), seq_id, flags);
if (n != data_tgt.size()) {
GGML_ABORT("checkpoint size mismatch: expected %zu, got %zu\n", data_tgt.size(), n);
}
return n == data_tgt.size();
}
void common_prompt_checkpoint::load_dft(
bool common_prompt_checkpoint::load_dft(
llama_context * ctx,
llama_seq_id seq_id,
llama_state_seq_flags flags) const {
if (ctx == nullptr) {
return;
return true;
}
if (data_dft.empty()) {
return;
return true;
}
const size_t n = llama_state_seq_set_data_ext(ctx, data_dft.data(), data_dft.size(), seq_id, flags);
if (n != data_dft.size()) {
GGML_ABORT("checkpoint size mismatch: expected %zu, got %zu\n", data_dft.size(), n);
}
return n == data_dft.size();
}
void common_prompt_checkpoint::clear_tgt() {
+7 -4
View File
@@ -593,7 +593,7 @@ struct common_params {
ggml_type cache_type_k = GGML_TYPE_F16; // KV cache data type for the K
ggml_type cache_type_v = GGML_TYPE_F16; // KV cache data type for the V
size_t moe_cache_size = 0; // GPU cache size in bytes for the MoE experts kept in the CPU
size_t moe_cache_size = 0; // GPU cache size in bytes for the MoE experts kept in the CPU, split among the GPUs like the layers
common_conversation_mode conversation_mode = COMMON_CONVERSATION_MODE_AUTO;
@@ -625,7 +625,7 @@ struct common_params {
std::string cls_sep = "\t"; // separator of classification sequences
// server params
int32_t port = 8080; // server listens on this network port
int32_t port = 9931; // server listens on this network port
bool reuse_port = false; // allow multiple sockets to bind to the same port
int32_t timeout_read = 3600; // http read timeout in seconds
int32_t timeout_write = timeout_read; // http write timeout in seconds
@@ -963,6 +963,8 @@ enum common_decision_type {
COMMON_DECISION_TYPE_LAYA, // score of one marker token per option, read from the embeddings output
COMMON_DECISION_TYPE_CLEF, // all questions in one prompt, score of option i read from the embeddings output at row i
COMMON_DECISION_TYPE_PPLX_DECIDER, // same as openjev, label codes of 1 or 2 letters
COMMON_DECISION_TYPE_LFM2_D1, // same as openjev, the labels depend on the question type
COMMON_DECISION_TYPE_LFM2_D1_OMNI, // same as laya, other prompt layout
COMMON_DECISION_TYPE_UNKNOWN, // a decision model of a type that is not supported
};
@@ -1294,12 +1296,13 @@ struct common_prompt_checkpoint {
llama_seq_id seq_id,
llama_state_seq_flags flags);
void load_tgt(
// return false if the state could not be restored
bool load_tgt(
llama_context * ctx,
llama_seq_id seq_id,
llama_state_seq_flags flags) const;
void load_dft(
bool load_dft(
llama_context * ctx,
llama_seq_id seq_id,
llama_state_seq_flags flags) const;
+5 -5
View File
@@ -152,13 +152,13 @@ common_chat_params common_chat_params_init_deepseek_v3_2(const common_chat_templ
// build tool call section first since we might need it in reasoning
auto tool_choice = p.choice();
if (has_tool_calls) {
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & function = tool.at("function");
std::string name = function.at("name");
std::vector<common_peg_parser> required_parsers;
std::vector<common_peg_parser> optional_parsers;
foreach_parameter(function, [&](const common_chat_schema_property & param, const common_chat_schema_document_ptr & doc) {
foreach_parameter(function, [&](size_t param_index, const common_chat_schema_property & param, const common_chat_schema_document_ptr & doc) {
bool is_string = param.schema->may_be_string();
auto arg = p.tool_arg(
@@ -166,11 +166,11 @@ common_chat_params common_chat_params_init_deepseek_v3_2(const common_chat_templ
p.literal("\" string=\"" + std::string(is_string ? "true" : "false") + "\">")) +
(is_string ?
p.tool_arg_string_value(p.until(PARAM_END)) :
p.tool_arg_json_value(p.schema(p.json(), "tool-" + name + "-arg-" + param.name + "-schema",
p.tool_arg_json_value(p.schema(p.json(), "tool-" + std::to_string(tool_index) + "-arg-" + std::to_string(param_index) + "-schema",
doc, *param.schema))) +
p.tool_arg_close(p.literal(PARAM_END)));
auto named_arg = p.rule("tool-" + name + "-arg-" + param.name, arg);
auto named_arg = p.rule("tool-" + std::to_string(tool_index) + "-arg-" + std::to_string(param_index), arg);
if (param.required) {
required_parsers.push_back(named_arg);
} else {
@@ -199,7 +199,7 @@ common_chat_params common_chat_params_init_deepseek_v3_2(const common_chat_templ
p.tool_name(p.literal(name)) + p.literal("\">\n")) +
invoke_body + p.space() + p.tool_close(p.literal(INVOKE_END)));
tool_choice |= p.rule("tool-" + name, func_parser);
tool_choice |= p.rule("tool-" + std::to_string(tool_index), func_parser);
});
}
+3 -3
View File
@@ -42,7 +42,7 @@ common_chat_params common_chat_params_init_functionary_v3_2(const common_chat_te
// Build tool call parsers for each available function
auto tool_choice = p.choice();
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & function = tool.at("function");
std::string name = function.at("name");
const auto schema = common_chat_tool_parameters(function);
@@ -50,10 +50,10 @@ common_chat_params common_chat_params_init_functionary_v3_2(const common_chat_te
// Tool format: >>>function_name\n{json_args}
auto tool_parser = p.tool(
p.tool_open(p.tool_name(p.literal(name)) + p.literal("\n")) +
p.tool_args(p.schema(p.json(), "tool-" + name + "-schema", schema))
p.tool_args(p.schema(p.json(), "tool-" + std::to_string(tool_index) + "-schema", schema))
);
tool_choice |= p.rule("tool-" + name, tool_parser);
tool_choice |= p.rule("tool-" + std::to_string(tool_index), tool_parser);
});
auto content_only = content_until_end;
+2 -2
View File
@@ -254,13 +254,13 @@ common_chat_params common_chat_params_init_gemma4(const common_chat_template &
auto tool_choice = p.choice();
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & function = tool.at("function");
std::string name = function.at("name");
// TODO @aldehir : need to extend json-schema-to-grammar to produce more than JSON rules
// const auto & params = function.at("parameters");
tool_choice |= p.rule("tool-" + name, p.tool(p.sequence({
tool_choice |= p.rule("tool-" + std::to_string(tool_index), p.tool(p.sequence({
p.tool_open(p.tool_name(p.literal(name)) + p.peek(p.literal("{"))),
p.tool_args(p.ref("gemma4-dict")),
})));
+4 -3
View File
@@ -30,17 +30,18 @@ common_chat_params common_chat_params_init_gigachat_v3(
if (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE) {
// Build a choice of all available tools
auto tool_choice = p.choice();
for (const auto & tool : inputs.tools) {
for (size_t i = 0; i < inputs.tools.size(); i++) {
const auto & tool = inputs.tools[i];
const auto & function = tool.at("function");
std::string name = function.at("name");
const auto schema = common_chat_tool_parameters(function);
auto tool_name = p.json_member("name", "\"" + p.tool_name(p.literal(name)) + "\"");
auto tool_args = p.json_member("arguments", p.tool_args(p.schema(p.json(), "tool-" + name + "-schema", schema)));
auto tool_args = p.json_member("arguments", p.tool_args(p.schema(p.json(), "tool-" + std::to_string(i) + "-schema", schema)));
auto tool_open = p.tool_open(p.literal("{") << tool_name);
tool_choice |= p.rule("tool-" + name, tool_open << "," << tool_args << "}");
tool_choice |= p.rule("tool-" + std::to_string(i), tool_open << "," << tool_args << "}");
}
// Define the tool call structure
+3 -3
View File
@@ -106,14 +106,14 @@ common_chat_params common_chat_params_init_gpt_oss(const common_chat_template &
if (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE) {
auto tool_choice = p.choice();
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & function = tool.at("function");
std::string name = function.at("name");
const auto params = common_chat_tool_parameters(function);
auto func_name = p.literal(" to=functions.") + p.tool_name(p.literal(name));
auto constraint = p.optional(p.space() + p.optional(p.literal("<|constrain|>")) + constrain_type);
auto args = p.tool_args(p.schema(p.json(), "tool-" + name + "-schema", params));
auto args = p.tool_args(p.schema(p.json(), "tool-" + std::to_string(tool_index) + "-schema", params));
// recipient in role header
// <|start|>assistant to=functions.NAME<|channel|>(commentary|analysis)[constraint]<|message|>ARGS
@@ -123,7 +123,7 @@ common_chat_params common_chat_params_init_gpt_oss(const common_chat_template &
// <|channel|>(commentary|analysis) to=functions.NAME[constraint]<|message|>ARGS
auto tool_in_channel = p.tool(p.tool_open(channel + func_name + constraint + p.literal("<|message|>")) + args);
tool_choice |= p.rule("tool-" + name, tool_in_role | tool_in_channel);
tool_choice |= p.rule("tool-" + std::to_string(tool_index), tool_in_role | tool_in_channel);
});
auto tool_call = p.trigger_rule("tool-call", tool_choice);
+5 -5
View File
@@ -110,14 +110,14 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template
// The models leave out <ifm|arg_type> even when asked for xml_typed
auto arg_type = call_format == "xml_typed" ? p.optional(ARG_TYPE + p.until(ARG_TYPE_END) + ARG_TYPE_END + p.space()) : p.eps();
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & function = tool.at("function");
std::string name = function.at("name");
std::vector<common_peg_parser> required_args;
std::vector<common_peg_parser> optional_args;
foreach_parameter(function, [&](const common_chat_schema_property & param, const common_chat_schema_document_ptr & doc) {
auto rule_name = "tool-" + name + "-arg-" + param.name;
foreach_parameter(function, [&](size_t param_index, const common_chat_schema_property & param, const common_chat_schema_document_ptr & doc) {
auto rule_name = "tool-" + std::to_string(tool_index) + "-arg-" + std::to_string(param_index);
auto types = param.schema->value_types();
auto arg_value = arg_string;
if (!types.has(common_chat_schema::TYPE_STRING)) {
@@ -149,12 +149,12 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template
(param.required ? required_args : optional_args).push_back(p.rule(rule_name, arg));
});
auto args = p.permute("tool-" + name + "-args", required_args);
auto args = p.permute("tool-" + std::to_string(tool_index) + "-args", required_args);
if (!optional_args.empty()) {
args = args + p.zero_or_more(p.choice(optional_args));
}
tool_choice |= p.rule("tool-" + name, p.tool(
tool_choice |= p.rule("tool-" + std::to_string(tool_index), p.tool(
p.tool_open(CALL_START + p.tool_name(p.literal(name)) + "\n") + p.tool_args(args) << p.tool_close(p.literal(CALL_END))));
});
}
+3 -3
View File
@@ -79,7 +79,7 @@ common_chat_params common_chat_params_init_kimi_k2(const common_chat_template &
// The ID format is: functions.<name>:<index>
// We need to match: functions.<name>:<digits>
auto tool_choice = p.choice();
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & function = tool.at("function");
std::string name = function.at("name");
const auto schema = common_chat_tool_parameters(function);
@@ -89,11 +89,11 @@ common_chat_params common_chat_params_init_kimi_k2(const common_chat_template &
auto tool_id = p.tool_id(p.literal("functions.") + p.tool_name(p.literal(name)) + p.literal(":") + p.chars("[0-9]", 1, -1));
auto tool_parser = p.tool(
p.tool_open(tool_id + p.literal(ARGS_BEGIN)) +
p.tool_args(p.schema(p.json(), "tool-" + name + "-schema", schema)) +
p.tool_args(p.schema(p.json(), "tool-" + std::to_string(tool_index) + "-schema", schema)) +
p.tool_close(p.optional((p.literal(CALL_END))))
);
tool_choice |= p.rule("tool-" + name, tool_parser);
tool_choice |= p.rule("tool-" + std::to_string(tool_index), tool_parser);
});
// Tool calls section: <|tool_calls_section_begin|> tool_calls <|tool_calls_section_end|>
+4 -3
View File
@@ -95,7 +95,7 @@ common_chat_params common_chat_params_init_kimi_k3(const common_chat_template &
}
auto tool_choices = p.choice();
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & function = tool.at("function");
std::string name = function.at("name");
const json schema = common_chat_tool_parameters(function);
@@ -106,6 +106,7 @@ common_chat_params common_chat_params_init_kimi_k3(const common_chat_template &
auto args = p.eps();
if (schema.contains("properties") && !schema.at("properties").empty()) {
auto arg_choices = p.choice();
size_t param_index = 0;
for (const auto & prop : schema.at("properties").items()) {
const std::string & key = prop.key();
@@ -119,7 +120,7 @@ common_chat_params common_chat_params_init_kimi_k3(const common_chat_template &
p.tool_arg_value(p.until(ARG_END));
// skip the trailing type="..." attribute: anything up to <|sep|>
arg_choices |= p.rule("kimi-k3-arg-" + name + "-" + key,
arg_choices |= p.rule("kimi-k3-arg-" + std::to_string(tool_index) + "-" + std::to_string(param_index++),
p.tool_arg(p.tool_arg_open(p.literal(ARG_START)) +
p.tool_arg_name(p.literal(key)) + p.literal("\"") +
p.until(SEP) + p.literal(SEP) + value +
@@ -133,7 +134,7 @@ common_chat_params common_chat_params_init_kimi_k3(const common_chat_template &
p.until(SEP) + p.literal(SEP)) +
p.tool_args(args) + p.tool_close(p.literal(CALL_END)));
tool_choices |= p.rule("kimi-k3-tool-" + name, call);
tool_choices |= p.rule("kimi-k3-tool-" + std::to_string(tool_index), call);
});
// all calls go inside one tools section, then the message is closed. the
+5 -5
View File
@@ -118,7 +118,7 @@ common_chat_params common_chat_params_init_ling3(const common_chat_template &
auto arg_string = p.rule("ling3-arg-string",
p.tool_arg_string_value(p.until(ARG_VAL_END)) + arg_close);
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & function = tool.at("function");
std::string name = function.at("name");
@@ -127,8 +127,8 @@ common_chat_params common_chat_params_init_ling3(const common_chat_template &
// each argument may be preceded by whitespace: the model emits
// newlines between arguments, the template history does not
foreach_parameter(function, [&](const common_chat_schema_property & param, const common_chat_schema_document_ptr & doc) {
auto rule_name = "ling3-arg-" + name + "-" + param.name;
foreach_parameter(function, [&](size_t param_index, const common_chat_schema_property & param, const common_chat_schema_document_ptr & doc) {
auto rule_name = "ling3-arg-" + std::to_string(tool_index) + "-" + std::to_string(param_index);
auto types = param.schema->value_types();
@@ -159,7 +159,7 @@ common_chat_params common_chat_params_init_ling3(const common_chat_template &
// required arguments in any order (as Qwen3-Coder does), then
// optional ones in any order and number
auto args = p.permute("ling3-" + name + "-args", required_args);
auto args = p.permute("ling3-" + std::to_string(tool_index) + "-args", required_args);
if (!optional_args.empty()) {
args = args + p.zero_or_more(p.choice(optional_args));
}
@@ -169,7 +169,7 @@ common_chat_params common_chat_params_init_ling3(const common_chat_template &
p.tool_args(args) +
p.tool_close(p.optional(p.space()) + p.literal(CALL_END)));
tool_choices |= p.rule("ling3-tool-" + name, call);
tool_choices |= p.rule("ling3-tool-" + std::to_string(tool_index), call);
});
auto calls = inputs.parallel_tool_calls ?
+3 -3
View File
@@ -109,13 +109,13 @@ common_chat_params common_chat_params_init_llm_jp_harmony(const common_chat_temp
if (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE) {
auto tool_choice = p.choice();
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & function = tool.at("function");
std::string name = function.at("name");
const auto params = common_chat_tool_parameters(function);
auto func_name = p.literal(" to=functions.") + p.tool_name(p.literal(name));
auto args = p.tool_args(p.schema(p.json(), "tool-" + name + "-schema", params));
auto args = p.tool_args(p.schema(p.json(), "tool-" + std::to_string(tool_index) + "-schema", params));
// recipient in role header
// <|start|>assistant to=functions.NAME<|channel|>(commentary|analysis)[constraint]<|message|>ARGS
@@ -125,7 +125,7 @@ common_chat_params common_chat_params_init_llm_jp_harmony(const common_chat_temp
// <|channel|>(commentary|analysis) to=functions.NAME[constraint]<|message|>ARGS
auto tool_in_channel = p.tool(p.tool_open(channel + func_name + constraint + message) + args);
tool_choice |= p.rule("tool-" + name, tool_in_role | tool_in_channel);
tool_choice |= p.rule("tool-" + std::to_string(tool_index), tool_in_role | tool_in_channel);
});
// parallel calls are separated by <|end|>; inside the trigger rule so the lazy grammar covers all of them
+4 -4
View File
@@ -68,18 +68,18 @@ common_chat_params common_chat_params_init_minicpm5(const common_chat_template &
});
auto tool_choice = p.choice();
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & function = tool.at("function");
const std::string name = function.at("name");
std::vector<common_peg_parser> arg_rules;
foreach_parameter(function, [&](const common_chat_schema_property & prop, const common_chat_schema_document_ptr & doc) {
foreach_parameter(function, [&](size_t param_index, const common_chat_schema_property & prop, const common_chat_schema_document_ptr & doc) {
auto value_parser = p.eps();
if (prop.schema->may_be_string()) {
value_parser = string_value;
} else {
value_parser = p.tool_arg_json_value(
p.schema(p.json(), "tool-" + name + "-arg-" + prop.name + "-schema", doc, *prop.schema)
p.schema(p.json(), "tool-" + std::to_string(tool_index) + "-arg-" + std::to_string(param_index) + "-schema", doc, *prop.schema)
) + p.tool_arg_close(p.literal("</param>"));
}
@@ -99,7 +99,7 @@ common_chat_params common_chat_params_init_minicpm5(const common_chat_template &
<< p.tool_args(args)
<< p.tool_close(p.literal("</function>")));
tool_choice |= p.rule("tool-" + name, tool_parser);
tool_choice |= p.rule("tool-" + std::to_string(tool_index), tool_parser);
});
auto max_calls = inputs.parallel_tool_calls ? -1 : 1;
+6 -5
View File
@@ -85,7 +85,7 @@ common_chat_params common_chat_params_init_minimax_m3(const common_chat_template
}
auto tool_choice = p.choice();
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & function = tool.at("function");
std::string name = function.at("name");
auto params = common_chat_tool_parameters(function);
@@ -154,8 +154,9 @@ common_chat_params common_chat_params_init_minimax_m3(const common_chat_template
members_of = [&](const common_chat_schema_object & object, const std::string & rule_prefix) -> common_peg_parser {
std::vector<common_peg_parser> required_elements;
std::vector<common_peg_parser> optional_elements;
for (const auto & prop : object.properties) {
auto element = element_of(prop.name, *prop.schema, rule_prefix + "-" + prop.name);
for (size_t i = 0; i < object.properties.size(); i++) {
const auto & prop = object.properties[i];
auto element = element_of(prop.name, *prop.schema, rule_prefix + "-" + std::to_string(i));
(prop.required ? required_elements : optional_elements).push_back(element);
}
@@ -180,7 +181,7 @@ common_chat_params common_chat_params_init_minimax_m3(const common_chat_template
common_peg_parser invoke_body = p.eps();
if (doc->root->kind() == common_chat_schema::KIND_OBJECT) {
invoke_body = members_of(static_cast<const common_chat_schema_object &>(*doc->root), "tool-" + name + "-arg");
invoke_body = members_of(static_cast<const common_chat_schema_object &>(*doc->root), "tool-" + std::to_string(tool_index) + "-arg");
}
auto func_parser = p.tool(
@@ -189,7 +190,7 @@ common_chat_params common_chat_params_init_minimax_m3(const common_chat_template
p.space() + invoke_body + p.space() +
p.tool_close(p.literal(INVOKE_END)));
tool_choice |= p.rule("tool-" + name, func_parser);
tool_choice |= p.rule("tool-" + std::to_string(tool_index), func_parser);
});
auto require_tools = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED;
+3 -3
View File
@@ -86,14 +86,14 @@ common_chat_params common_chat_params_init_ministral_3(const common_chat_templat
// Tool call parser
if (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE) {
auto tool_choice = p.choice();
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & function = tool.at("function");
std::string name = function.at("name");
const auto schema = common_chat_tool_parameters(function);
tool_choice |=
p.rule("tool-" + name, p.tool_open(p.tool_name(p.literal(name)) + "[ARGS]") +
p.tool_args(p.schema(p.json(), "tool-" + name + "-schema", schema)));
p.rule("tool-" + std::to_string(tool_index), p.tool_open(p.tool_name(p.literal(name)) + "[ARGS]") +
p.tool_args(p.schema(p.json(), "tool-" + std::to_string(tool_index) + "-schema", schema)));
});
auto min_calls = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED ? 1 : 0;
+4 -4
View File
@@ -81,18 +81,18 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
"</atem:parameter>");
auto tool_choice = p.choice();
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & function = tool.at("function");
const std::string name = function.at("name");
std::vector<common_peg_parser> arg_rules;
foreach_parameter(function, [&](const common_chat_schema_property & prop, const common_chat_schema_document_ptr & doc) {
foreach_parameter(function, [&](size_t param_index, const common_chat_schema_property & prop, const common_chat_schema_document_ptr & doc) {
auto value_parser = p.eps();
if (prop.schema->may_be_string()) {
value_parser = string_value;
} else {
value_parser = p.tool_arg_json_value(
p.schema(p.json(), "tool-" + name + "-arg-" + prop.name + "-schema", doc, *prop.schema))
p.schema(p.json(), "tool-" + std::to_string(tool_index) + "-arg-" + std::to_string(param_index) + "-schema", doc, *prop.schema))
+ p.tool_arg_close(p.literal("</atem:parameter>"));
}
@@ -113,7 +113,7 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
<< p.tool_args(args)
<< p.tool_close(p.literal("</atem:invoke>") + p.space() + p.literal("</atem:function_calls>")));
tool_choice |= p.rule("tool-" + name, tool_parser);
tool_choice |= p.rule("tool-" + std::to_string(tool_index), tool_parser);
});
auto tool_calls = inputs.parallel_tool_calls
+7 -6
View File
@@ -2,24 +2,25 @@
#include "log.h"
void foreach_function(const json & tools, const std::function<void(const json &)> & fn) {
for (const auto & tool : tools) {
void foreach_function(const json & tools, const std::function<void(size_t, const json &)> & fn) {
for (size_t i = 0; i < tools.size(); i++) {
const auto & tool = tools[i];
if (!tool.contains("type") || tool.at("type") != "function" || !tool.contains("function")) {
LOG_INF("Skipping tool without function: %s", tool.dump(2).c_str());
continue;
}
fn(tool);
fn(i, tool);
}
}
void foreach_parameter(const json & function, const std::function<void(const common_chat_schema_property &, const common_chat_schema_document_ptr &)> & fn) {
void foreach_parameter(const json & function, const std::function<void(size_t, const common_chat_schema_property &, const common_chat_schema_document_ptr &)> & fn) {
auto params = common_chat_tool_parameters(function);
auto doc = std::make_shared<const common_chat_schema_document>(common_chat_schema_from_json(params));
const auto * object = dynamic_cast<const common_chat_schema_object *>(doc->root.get());
if (!object) {
return;
}
for (const auto & prop : object->properties) {
fn(prop, doc);
for (size_t i = 0; i < object->properties.size(); i++) {
fn(i, object->properties[i], doc);
}
}
+6 -4
View File
@@ -17,11 +17,11 @@
using json = common_json;
// iterate over the function tools of an OpenAI-style tools array
void foreach_function(const json & tools, const std::function<void(const json &)> & fn);
// iterate over the function tools of an OpenAI-style tools array, passing each tool with its index in the array
void foreach_function(const json & tools, const std::function<void(size_t, const json &)> & fn);
// iterate over the parameters of a function tool, with the document that owns them
void foreach_parameter(const json & function, const std::function<void(const common_chat_schema_property &, const common_chat_schema_document_ptr &)> & fn);
// iterate over the parameters of a function tool, passing each parameter with its index and the document that owns it
void foreach_parameter(const json & function, const std::function<void(size_t, const common_chat_schema_property &, const common_chat_schema_document_ptr &)> & fn);
// render a template; the override arguments let a parser feed in messages, tools or context it has rewritten
std::string common_chat_template_direct_apply_impl(
@@ -81,3 +81,5 @@ common_chat_params common_chat_params_init_ministral_3(const common_chat_templat
common_chat_params common_chat_params_init_muse_glimmer(const common_chat_template & tmpl, const autoparser::generation_params & inputs);
common_chat_params common_chat_params_init_qwen3_coder(const common_chat_template & tmpl, const autoparser::generation_params & inputs);
common_chat_params common_chat_params_init_translate_gemma(const common_chat_template & tmpl, const autoparser::generation_params & inputs);
+6 -6
View File
@@ -65,7 +65,7 @@ common_chat_params common_chat_params_init_qwen3_coder(const common_chat_templat
// Match complete <function=name> opener for Qwen3-Coder models that occasionally omit the
// starting <tool_call>. The model may hallucinate a tool name, but it is preferable over
// constraining on <function which may occur in valid content generation, e.g. #include <functional>
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t, const json & tool) {
const std::string name = tool.at("function").at("name");
tool_call_starts.push_back("<function=" + name + ">");
});
@@ -93,15 +93,15 @@ common_chat_params common_chat_params_init_qwen3_coder(const common_chat_templat
p.ac(p.tool_arg_string_value(p.until("\n</parameter>\n")) + arg_close, "\n</parameter>\n"));
auto tool_choice = p.choice();
foreach_function(inputs.tools, [&](const json & tool) {
foreach_function(inputs.tools, [&](size_t tool_index, const json & tool) {
const auto & function = tool.at("function");
std::string name = function.at("name");
std::vector<common_peg_parser> required_args;
std::vector<common_peg_parser> optional_args;
foreach_parameter(function, [&](const common_chat_schema_property & param, const common_chat_schema_document_ptr & doc) {
auto rule_name = "tool-" + name + "-arg-" + param.name;
foreach_parameter(function, [&](size_t param_index, const common_chat_schema_property & param, const common_chat_schema_document_ptr & doc) {
auto rule_name = "tool-" + std::to_string(tool_index) + "-arg-" + std::to_string(param_index);
auto arg_open = p.tool_arg_open("<parameter=" + p.tool_arg_name(p.literal(param.name)) + ">\n");
@@ -141,7 +141,7 @@ common_chat_params common_chat_params_init_qwen3_coder(const common_chat_templat
// Accept required arguments in any order, as Qwen does not always adhere to the
// order provided.
auto args = p.permute("tool-" + name + "-args", required_args);
auto args = p.permute("tool-" + std::to_string(tool_index) + "-args", required_args);
if (!optional_args.empty()) {
args = args + p.zero_or_more(p.choice(optional_args));
}
@@ -150,7 +150,7 @@ common_chat_params common_chat_params_init_qwen3_coder(const common_chat_templat
p.tool_args(args) +
p.tool_close(p.literal("</function>\n")));
tool_choice |= p.rule("tool-" + name, func);
tool_choice |= p.rule("tool-" + std::to_string(tool_index), func);
});
auto min_calls = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED ? 1 : 0;
+1
View File
@@ -20,4 +20,5 @@ set(LLAMA_CHAT_PARSERS_SOURCES
${CMAKE_CURRENT_LIST_DIR}/ministral3.cpp
${CMAKE_CURRENT_LIST_DIR}/muse-glimmer.cpp
${CMAKE_CURRENT_LIST_DIR}/qwen3-coder.cpp
${CMAKE_CURRENT_LIST_DIR}/translate-gemma.cpp
)
+63
View File
@@ -0,0 +1,63 @@
#include "parsers.h"
#include "log.h"
// TranslateGemma does not support tools or reasoning, it only needs user messages in its own content schema
common_chat_params common_chat_params_init_translate_gemma(
const common_chat_template & tmpl,
const autoparser::generation_params & inputs) {
common_chat_params data;
// default to chat_template_kwargs, or en-GB if not specified
std::string src_lang = inputs.extra_context.value("source_lang_code", "en-GB");
std::string tgt_lang = inputs.extra_context.value("target_lang_code", "en-GB");
for (const char * key : { "source_lang_code", "target_lang_code" }) {
if (!inputs.extra_context.contains(key)) {
LOG_WRN("TranslateGemma: %s not set in chat_template_kwargs, defaulting to en-GB\n", key);
}
}
json messages = inputs.messages;
for (auto & message : messages) {
if (message.value("role", "") != "user") {
continue;
}
std::string text;
const auto & content = message.contains("content") ? message.at("content") : json();
if (content.is_string()) {
text = content.get<std::string>();
} else if (content.is_array()) {
for (const auto & part : content) {
if (!text.empty()) {
text += "\n";
}
text += part.value("text", "");
}
}
message["content"] = json::array({
json{
{"type", "text"},
{"text", text},
{"source_lang_code", src_lang},
{"target_lang_code", tgt_lang},
}
});
}
data.prompt = common_chat_template_direct_apply_impl(tmpl, inputs, messages);
data.generation_prompt = common_chat_template_generation_prompt_impl(tmpl, inputs, messages);
data.format = COMMON_CHAT_FORMAT_PEG_NATIVE;
data.supports_thinking = false;
if (inputs.has_continuation()) {
data.generation_prompt = "<start_of_turn>model\n" + inputs.continue_msg.render_content();
data.prompt += data.generation_prompt;
}
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
return p.literal(data.generation_prompt) << p.content(p.rest());
});
data.parser = parser.save();
return data;
}
+4
View File
@@ -161,6 +161,8 @@ TEXT_MODEL_MAP: dict[str, str] = {
"Lfm2BidirectionalForMaskedLM": "lfm2",
"Lfm2BidirectionalModel": "lfm2",
"Lfm2ForCausalLM": "lfm2",
"D1Model": "lfm2",
"D1OmniModel": "lfm2",
"Lfm2Model": "lfm2",
"Lfm2MoeForCausalLM": "lfm2",
"Llama4ForCausalLM": "llama",
@@ -253,6 +255,7 @@ TEXT_MODEL_MAP: dict[str, str] = {
"Qwen3_5ForConditionalGeneration": "qwen",
"Qwen3_5MoeForCausalLM": "qwen",
"Qwen3_5MoeForConditionalGeneration": "qwen",
"Qwen3_5TextModel": "qwen",
"Qwen4ExpForCausalLM": "qwen4exp",
"Qwen4ExpForConditionalGeneration": "qwen4exp",
"RND1": "qwen",
@@ -334,6 +337,7 @@ MMPROJ_MODEL_MAP: dict[str, str] = {
"KimiK25ForConditionalGeneration": "kimivl",
"KimiVLForConditionalGeneration": "kimivl",
"Lfm2AudioForConditionalGeneration": "lfm2",
"D1OmniModel": "lfm2",
"Lfm2VlForConditionalGeneration": "lfm2",
"LightOnOCRForConditionalGeneration": "lighton_ocr",
"Llama4ForConditionalGeneration": "llama4",
+14
View File
@@ -2336,12 +2336,26 @@ class TextModel(ModelBase):
else:
raise NotImplementedError("Only MEAN, CLS, and LAST pooling types supported")
self.gguf_writer.add_pooling_type(pooling_type)
else:
embedding_config_path = self.dir_model / "embedding_config.json"
if embedding_config_path.is_file():
with open(embedding_config_path, encoding="utf-8") as f:
embedding_config = json.load(f)
pooling = embedding_config.get("pooling")
if pooling == "last_token":
self.gguf_writer.add_pooling_type(gguf.PoolingType.LAST)
elif pooling is not None:
raise NotImplementedError(f"unsupported embedding_config.json pooling {pooling!r}")
# pooling before a classification head (e.g. ModernBertForSequenceClassification)
if (classifier_pooling := self.hparams.get("classifier_pooling")) is not None:
if classifier_pooling not in ("cls", "mean"):
raise NotImplementedError(f"Unsupported classifier_pooling: {classifier_pooling}")
self.gguf_writer.add_classifier_pooling_type(mode_mapping[classifier_pooling])
if (classifier_activation := self.hparams.get("classifier_activation")) is not None:
if classifier_activation not in ("gelu", "silu", "tanh"):
raise NotImplementedError(f"Unsupported classifier_activation: {classifier_activation}")
self.gguf_writer.add_classifier_activation(classifier_activation)
def _set_vocab_glmedge(self):
from transformers import AutoTokenizer
+2 -2
View File
@@ -849,8 +849,8 @@ class Gemma4DSparkModel(DFlashModel):
raise ValueError("Gemma4 DSpark attention bias and MoE are not supported")
if (self.hparams.get("draft_vocab_size") or self.hparams["vocab_size"]) != self.hparams["vocab_size"]:
raise ValueError("Gemma4 DSpark currently requires a full draft vocabulary")
if "model.lm_head.weight" not in self.model_tensors and self.hparams.get("tie_word_embeddings") is not True:
raise ValueError("Gemma4 DSpark requires lm_head.weight unless tie_word_embeddings is true")
if "model.lm_head.weight" not in self.model_tensors:
raise ValueError("Gemma4 DSpark requires lm_head.weight")
self.dflash_config = self.hparams.get("dflash_config", {})
markov_type = self.dflash_config.get("markov_head_type", self.hparams.get("markov_head_type", "vanilla"))
+239 -1
View File
@@ -1,5 +1,8 @@
from __future__ import annotations
import json
from pathlib import Path
from typing import Any, Callable, Iterable, TYPE_CHECKING
import torch
@@ -7,7 +10,7 @@ import torch
if TYPE_CHECKING:
from torch import Tensor
from .base import MmprojModel, ModelBase, TextModel, gguf
from .base import MmprojModel, ModelBase, TextModel, gguf, jinja_str_or_json, logger
from .gemma import ConformerAudioModel
@@ -65,6 +68,68 @@ class LFM2Model(TextModel):
yield from super().modify_tensors(data_torch, name, bid)
def _is_d1_checkpoint(dir_model: Path) -> bool:
if not (dir_model / "config.json").is_file():
return False
with open(dir_model / "config.json", encoding="utf-8") as f:
return json.load(f).get("auto_map", {}).get("AutoModel", "").endswith(".D1Model")
@ModelBase.register_hparams_loader(_is_d1_checkpoint)
def _load_d1_hparams(dir_model: Path) -> dict[str, Any]:
logger.info("gguf: detected d1 checkpoint")
hparams = ModelBase.load_hparams(dir_model, False, guess=False)
# the mmproj stays LFM2-VL
hparams["text_config"]["architectures"] = ["D1Model"]
return hparams
@ModelBase.register("D1Model")
@ModelBase.example("LiquidAI/d1-3b")
class D1Model(LFM2Model):
model_arch = gguf.MODEL_ARCH.LFM2
def set_vocab(self):
super().set_vocab()
self.gguf_writer.add_chat_template([{"name": "systemone", "template": self._systemone_template()}])
@staticmethod
def _systemone_template() -> str:
# follows prompt.py of the model repo
description = jinja_str_or_json("o.description")
choice = (
"{{ '\\n\\nOptions:\\n' }}"
"{% for o in options %}{{ o.label }} {% if o.description %}" + description + "{% else %}{{ o.key | replace('_', ' ') }}{% endif %}"
"{% if not loop.last %}{{ '\\n' }}{% endif %}{% endfor %}"
"{{ '\\n\\nReply with the option code only.' }}"
)
# with criteria, a missing description is written as None
noul = (
"{% set ns = namespace(criteria=false) %}{% for o in options %}{% if o.description is not none %}{% set ns.criteria = true %}{% endif %}{% endfor %}"
"{% if ns.criteria %}"
"{% for o in options %}{{ '\\nYes: ' if o.key == 'true' else '\\nNo: ' }}"
"{% if o.description is none %}None{% else %}" + description + "{% endif %}{% endfor %}{% endif %}"
"{{ '\\n\\nReply with yes or no only.' }}"
)
score = (
"{{ '\\n\\n' }}{% for o in options %}{{ o.key }} " + description + "{{ '\\n' }}{% endfor %}"
"{{ '\\nReply with a single digit 0-' }}{{ options | length - 1 }}{{ ' only.' }}"
)
return (
"<|startoftext|><|im_start|>user\n"
"{% for image in images %}{{ image }}{% endfor %}"
"{% if state is not none %}{% if state is string %}{{ state }}{% else %}{{ state | tojson(indent=2) }}{% endif %}"
"{{ '\\n\\n\\nQUESTION:\\n' }}{% endif %}"
+ jinja_str_or_json("instructions")
+ "{% if type == 'choice' %}" + choice + "{% elif type == 'noul' %}" + noul + "{% else %}" + score + "{% endif %}"
"{{ '<|im_end|>\\n<|im_start|>assistant\\n' }}"
)
def set_gguf_parameters(self):
super().set_gguf_parameters()
self.gguf_writer.add_decision_type(gguf.DecisionType.LFM2_D1)
@ModelBase.register("Lfm2Model", "Lfm2BidirectionalModel", "Lfm2BidirectionalForMaskedLM")
@ModelBase.example("LiquidAI/LFM2.5-ColBERT-350M", "LiquidAI/LFM2.5-Embedding-350M", "LiquidAI/LFM2.5-Encoder-350M", "LiquidAI/LFM2.5-Encoder-230M")
class LFM2ColBertModel(LFM2Model):
@@ -96,6 +161,121 @@ class LFM2ColBertModel(LFM2Model):
yield f"{self.dense_tensor_name}.weight", tensor.clone()
def _is_d1_omni_checkpoint(dir_model: Path) -> bool:
if not (dir_model / "config.json").is_file():
return False
with open(dir_model / "config.json", encoding="utf-8") as f:
return json.load(f).get("model_type") == "d1_omni"
@ModelBase.register_hparams_loader(_is_d1_omni_checkpoint)
def _load_d1_omni_hparams(dir_model: Path) -> dict[str, Any]:
logger.info("gguf: detected d1-omni checkpoint")
hparams = ModelBase.load_hparams(dir_model, False, guess=False)
text = hparams["text_config"]
n_layer, n_layer_head = text["num_hidden_layers"], hparams["head_layers"]
# the trunk uses the LFM2 FFN sizing, the head blocks are appended with a plain 4x MLP
n_ff = int(text["block_ffn_dim_multiplier"] * int(2 * text["intermediate_size"] / 3))
n_ff = text["block_multiple_of"] * ((n_ff + text["block_multiple_of"] - 1) // text["block_multiple_of"])
text["num_hidden_layers"] = n_layer + n_layer_head
text["intermediate_size"] = [n_ff] * n_layer + [4 * text["hidden_size"]] * n_layer_head
text["block_auto_adjust_ff_dim"] = False
return hparams
@ModelBase.register("D1OmniModel")
@ModelBase.example("LiquidAI/d1-omni-600M")
class D1OmniModel(LFM2Model):
model_arch = gguf.MODEL_ARCH.LFM2
# the server cuts the text to these lengths, see server-decision.cpp
_MAX_LENGTH = 16384
_IMAGE_TEXT_LENGTH = 896
_AUDIO_TEXT_LENGTH = 15360
def set_vocab(self):
super().set_vocab()
# the systemone template writes the BOS, after the media
self.gguf_writer.remove_key(gguf.Keys.Tokenizer.ADD_BOS)
self.gguf_writer.add_add_bos_token(False)
self.gguf_writer.add_token_type_count(3) # choice, score, noul
self.gguf_writer.add_chat_template([{"name": "systemone", "template": self._systemone_template()}])
@staticmethod
def _systemone_template() -> str:
# follows prompt.py of the model repo, the server cuts each marked piece to its token budget
# the media (images, or an audio clip if audio is true) come first
description = jinja_str_or_json("o.description")
has_description = "o.description is not none and o.description != ''"
yes_no = "{{ 'yes' if o.key == 'true' else 'no' }}"
option_code = "{% if loop.index0 < 10 %}00{% elif loop.index0 < 100 %}0{% endif %}{{ loop.index0 }}"
option = (
"{% if type == 'choice' and audio %}option_" + option_code + ": "
"{% if " + has_description + " %}" + description + "{% else %}{{ o.key }}{% endif %}"
"{% elif type == 'choice' %}{{ o.key }}{% if " + has_description + " %}: " + description + "{% endif %}"
"{% elif type == 'score' %}level {{ o.key }}: " + description
+ "{% elif audio %}{{ o.key }}: " + yes_no
+ "{% else %}{{ o.key }}: {% if " + has_description + " %}" + description
+ "{% elif images and not ns.criteria %}" + yes_no
+ "{% elif o.key == 'true' %}yes, the statement holds"
"{% else %}no, the statement does not hold{% endif %}{% endif %}"
)
state = "{% if state is string %}{{ state }}{% elif state is not none %}{{ state | tojson }}{% elif audio %}{}{% endif %}"
return (
"{% set ns = namespace(criteria=false) %}"
"{% for o in options %}{% if o.description is not none %}{% set ns.criteria = true %}{% endif %}{% endfor %}"
"{% for image in images %}{{ image }}{% endfor %}{{ sep }}"
"<|startoftext|><|reserved_7|>{{ sep }}{{ mark_state }}" + state
+ "{{ sep }}{{ mark_question }}<|reserved_8|>" + jinja_str_or_json("instructions")
+ "{% for o in options %}{{ sep }}<|reserved_9|><|mask|>{{ sep }}{{ mark_option }} " + option
+ "{{ sep }}<|reserved_10|>{% endfor %}{{ sep }}<|reserved_11|>"
)
def set_gguf_parameters(self):
lengths = (self.hparams["max_length"], self.hparams["image_text_length"], self.hparams["audio_text_length"])
if lengths != (self._MAX_LENGTH, self._IMAGE_TEXT_LENGTH, self._AUDIO_TEXT_LENGTH):
raise ValueError(f"unexpected text lengths: {lengths}")
n_head, n_layer_head = self.hparams["num_attention_heads"], self.hparams["head_layers"]
self.hparams["num_key_value_heads"] = [
self.hparams["num_key_value_heads"] if t != "conv" else 0 for t in self.hparams["layer_types"]
] + [n_head] * n_layer_head
# the head needs per-layer sizes, LFM2Model writes a single feed forward length
TextModel.set_gguf_parameters(self)
self.gguf_writer.add_vocab_size(self.hparams["vocab_size"])
self.gguf_writer.add_shortconv_l_cache(self.hparams["conv_L_cache"])
self.gguf_writer.add_layer_norm_eps(1e-5) # nn.LayerNorm of the head
self.gguf_writer.add_causal_attention(False)
self.gguf_writer.add_decision_type(gguf.DecisionType.LFM2_D1_OMNI)
self.gguf_writer.add_decision_block_count(n_layer_head)
# "choice:3-5" -> "choice.3_5", "choice:11+" -> "choice.11"
for name, value in self.hparams["temperatures"].items():
self.gguf_writer.add_decision_temperature(name.replace(":", ".").replace("-", "_").rstrip("+"), value)
@classmethod
def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None:
name, gen = item
if name.startswith(("vision.", "audio.")):
return None
name = name.replace("encoder.", "model.", 1) if name.startswith("encoder.") else name
name = name.replace("head.head.layers.", "head.layers.").replace("in_proj_", "in_proj.")
name = name.removeprefix("head.") if name.startswith(("head.type_emb", "head.scorer")) else name
return super().filter_tensors((name, gen))
def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:
if name.startswith("head.layers.") and bid is not None:
# the head blocks come after the trunk blocks
suffix = name.split(".", 3)[3]
bid += self.block_count - self.hparams["head_layers"]
name = f"head.layers.{bid}.{suffix}"
yield from super().modify_tensors(data_torch, name, bid)
@ModelBase.register("Lfm2MoeForCausalLM")
@ModelBase.example("LiquidAI/LFM2-8B-A1B")
class LFM2MoeModel(TextModel):
@@ -188,6 +368,12 @@ class LFM2VLModel(MmprojModel):
# python notation, e.g. for vision_feature_layer == -1, we pick last layer -> vision_feature_layers_to_drop = 0
vision_feature_layers_to_drop = -(self.global_config.get("vision_feature_layer", -1) + 1)
self.gguf_writer.add_vision_block_count(self.find_vparam(self.n_block_keys) - vision_feature_layers_to_drop)
# PIL resample enum
if (resample := self.preprocessor_config.get("resample")) is not None:
resize_algo = {1: "lanczos", 2: "bilinear", 3: "bicubic"}.get(resample)
if resize_algo is None:
raise ValueError(f"unsupported resample: {resample}")
self.gguf_writer.add_vision_image_resize_algo(resize_algo)
@classmethod
def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None:
@@ -205,6 +391,58 @@ class LFM2VLModel(MmprojModel):
yield from super().modify_tensors(data_torch, name, bid)
@ModelBase.register("D1OmniModel")
@ModelBase.example("LiquidAI/d1-omni-600M")
class D1OmniMmprojModel(ConformerAudioModel):
has_vision_encoder = True
has_audio_encoder = True
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
assert self.hparams_vision is not None and self.hparams_audio is not None
# dynamic resolution, as LFM2VLModel
self.hparams_vision["image_size"] = 256
# the images are normalized to [-1, 1] (vision.py of the model repo)
self.preprocessor_config = {**self.preprocessor_config, "image_mean": [0.5] * 3, "image_std": [0.5] * 3}
self.hparams_audio["hidden_size"] = self.hparams_audio["d_model"]
self.hparams_audio["intermediate_size"] = self.hparams_audio["d_model"] * self.hparams_audio["ff_expansion_factor"]
self.hparams_audio["num_attention_heads"] = self.hparams_audio["n_heads"]
def set_gguf_parameters(self):
super().set_gguf_parameters()
self.gguf_writer.add_clip_vision_projector_type(gguf.VisionProjectorType.D1OMNI_V)
self.gguf_writer.add_vision_attention_layernorm_eps(self.find_vparam(["layer_norm_eps"]))
self.gguf_writer.add_vision_projector_scale_factor(self.global_config.get("downsample_factor", 2))
self.gguf_writer.add_vision_use_gelu(True)
assert self.hparams_audio is not None
self.gguf_writer.add_clip_audio_projector_type(gguf.VisionProjectorType.D1OMNI_A)
self.gguf_writer.add_audio_num_mel_bins(self.hparams_audio["feat_in"])
self.gguf_writer.add_audio_attention_layernorm_eps(1e-5)
@classmethod
def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None:
name, gen = item
if name.startswith(("encoder.", "head.")):
return None
name = name.replace("vision.tower.", "vision_tower.").replace("vision.projector.", "multi_modal_projector.")
name = name.replace("audio.encoder.", "conformer.")
# the residual block continues the adapter: norm, linear, gelu, linear, then norm, down, up
for old, new in (("adapter.norm", 0), ("adapter.linear_1", 1), ("adapter.linear_2", 3),
("residual.ln", 4), ("residual.down", 5), ("residual.up", 6)):
name = name.replace(f"audio.{old}.", f"audio_adapter.model.{new}.")
return super().filter_tensors((name, gen))
def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:
if "patch_embedding.weight" in name:
data_torch = data_torch.view(data_torch.shape[0], 16, 16, 3).permute(0, 3, 1, 2)
yield from super().modify_tensors(data_torch, name, bid)
@ModelBase.register("Lfm2AudioForConditionalGeneration")
@ModelBase.example("LiquidAI/LFM2.5-Audio-1.5B", "LiquidAI/LFM2-Audio-1.5B")
class LFM2AudioModel(ConformerAudioModel):
+14 -1
View File
@@ -650,11 +650,24 @@ class _Qwen35MRopeMixin:
self.gguf_writer.add_rope_dimension_sections(self._QWEN35_DEFAULT_MROPE_SECTION)
@ModelBase.register("Qwen3_5ForConditionalGeneration", "Qwen3_5ForCausalLM")
@ModelBase.register("Qwen3_5ForConditionalGeneration", "Qwen3_5ForCausalLM", "Qwen3_5TextModel")
@ModelBase.example("Qwen/Qwen3.5-9B")
class Qwen3_5TextModel(_Qwen35MRopeMixin, _LinearAttentionVReorderBase):
model_arch = gguf.MODEL_ARCH.QWEN35
def __init__(self, dir_model, *args, **kwargs):
# Inner TextModel does not own mtp.*. Set no_mtp before mixin bumps block_count.
hparams = kwargs.pop("hparams", None)
if hparams is None:
hparams = ModelBase.load_hparams(dir_model, self.is_mistral_format)
if get_model_architecture(hparams, ModelType.TEXT) == "Qwen3_5TextModel":
self.no_mtp = True
super().__init__(dir_model, *args, hparams=hparams, **kwargs)
def set_gguf_parameters(self):
super().set_gguf_parameters()
self._try_set_pooling_type()
def _is_openjev_checkpoint(dir_model: Path) -> bool:
return (dir_model / "helper" / "shim.py").is_file() and (dir_model / "config.json").is_file()
+5
View File
@@ -803,6 +803,7 @@ User can use the device management in [docs/multi-gpu.md](https://github.com/ggm
| GGML_SYCL_ENABLE_GRAPH | 0 (default) or 1 | Enable running computations through SYCL Graphs feature. Disabled by default because SYCL Graph is still on development, no better performance. |
| GGML_SYCL_ENABLE_HOST_PINNED_MEM | 0 or 1 (default) | Enable host pinned memory to speed up copy data from host to device. When disable it, host memory will common malloc() on CPU. Disable it when use `--load-model mlock`.|
| GGML_SYCL_HOST_PINNED_MEM_2G | 0 (default) or 1 | Limit the max memory allocation to be no more than 2GB when enable host pinned memory. USM allocations above 2 GiB take the relaxed/large-allocation path, which serializes H2D copies with compute and prevents copy/compute overlap. It will impact the startup time. Need more test. Depend on `GGML_SYCL_ENABLE_HOST_PINNED_MEM=1`.|
| GGML_SYCL_UPLOAD_STAGING_SLOTS | 4 (default) or non-negative integer | Number of 8 MiB pinned host slots used to stage tensor uploads (model loading), so the host copy of one slot overlaps the transfer of the previous one. Set to 0 to use the old path: a malloc'd bounce buffer and a blocking copy per tensor. |
| GGML_SYCL_GET_MEM_API | 0 (default) or 1 | Set to get memory info (free, total) by Level Zero or SYCL API:<br>0 - Level Zero API: support more GPUs, only run on Level Zero running time. When there is an error, fallback to call SYCL API. Depend on GGML_SYCL_SUPPORT_LEVEL_ZERO_API.<br>1 - SYCL API: legacy, support more running time, it can't get the free size of some GPUs (like Arc770). In such case, return the free size as value of total size.|
| GGML_SYCL_USE_LEVEL_ZERO_API | 1 (default) or 0 | Use Level Zero API for device memory allocation instead of SYCL. Reduces system RAM usage on Intel dGPUs by avoiding DMA-buf/TTM host memory staging. Requires GGML_SYCL_SUPPORT_LEVEL_ZERO_API=ON at build time. SYCL backend always runs on Level Zero running time even if it's set as OFF (The SYCL api will be usage for memory allocation).|
| GGML_SYCL_ENABLE_DNN | 0 or 1 (default)| Enable running computations through oneDNN and always use oneMKL. |
@@ -816,6 +817,10 @@ User can use the device management in [docs/multi-gpu.md](https://github.com/ggm
| GGML_SYCL_MKL_FA_DIAG | 0 (default) or 1 | Enable output fingerprinting for MKL flash attention. Dumps the first 64 float output values for the first 6 FA calls with n_kv ≥ 1024, labeled with kernel type (MKL/TILE/VEC) for cross-kernel comparison. |
| GGML_SYCL_ENABLE_FUSION | 0 or 1 (default) | Enable fused-kernel dispatch in graph compute. Unsupported types and layouts fall back to the standalone op kernels. See `ggml_sycl_can_fuse()`. |
| GGML_SYCL_ENABLE_ESIMD | 0 or 1 (default)| Enable ESIMD kernels when available. |
| GGML_SYCL_XMX_GATHER_TYPES | decimal bitmask, all bits set (default) | Weight formats that may use the XMX dequant-GEMM paths, which dequantize weights straight into the XMX tiles. This speeds up prompt processing of MoE models on GPUs with XMX units (Arc A- and B-series, Arc Pro, Data Center GPU Max), for example pp512 of Qwen3-30B-A3B UD-IQ3_XXS by about 50% on an Arc Pro B60. Bits:<br>* 1: IQ4_NL, 2: IQ3_S, 4: IQ4_XS, 8: IQ3_XXS, 16: IQ2_XXS, 32: IQ2_XS, 64: IQ2_S, 128: IQ1_S, 256: IQ1_M<br>* 512: Q8_0, 1024: Q4_K, 2048: Q5_K, 4096: Q6_K (MoE `MUL_MAT_ID` only)<br>Add values to combine them, for example `3` for IQ4_NL and IQ3_S; `0` disables the paths. A set bit does not force the path: batches of more than 64 tokens per expert or row lengths that are not a multiple of 256 (32 for IQ4_NL and Q8_0) use the library GEMM. |
| GGML_SYCL_XMX_GATHER_SHAPES | decimal bitmask, 255 (default) | XMX `joint_matrix` combinations the paths of `GGML_SYCL_XMX_GATHER_TYPES` may use; the operand type comes from `GGML_SYCL_DYNAMIC_PRECISION` and the best supported combination is picked automatically (logged as `fg_pick_combo`). Bits:<br>* Xe2, Xe3, Xe-HPC: 1: f16 8x16x16, 2: f16 16x16x16, 4: f16 32x64x16, 8: f16 32x64x32, 32: tf32 8x16x8, 64: bf16 8x16x16<br>* Xe-HPG (Arc A770, ARL-H): 16: f16 8x8x16, 128: bf16 8x8x16<br>Clear a bit to exclude a combination, or set a single bit to force one for testing. |
| GGML_SYCL_DYNAMIC_PRECISION | `F16` (default with `GGML_SYCL_F16=ON`), `BF16`, `TF32` or `F32` (default otherwise) | Operand type of the XMX dequant-GEMM paths (`GGML_SYCL_XMX_GATHER_TYPES`); accumulation is always f32. `F16` is the fastest, but activations above 65504 overflow. `BF16` keeps the f32 range at a 7-bit mantissa, `TF32` keeps the range and the f16 mantissa but is about 30% slower and needs Xe2, Xe3 or Xe-HPC, and `F32` turns the XMX paths off. Ops that request a higher src1 precision ([TAG_GGML_PREC]) get it regardless of this setting. |
| GGML_SYCL_DYNAMIC_REQUIRED_PRECISION | `F32` (default), `TF32`, `BF16` or `F16` | Lowest type the XMX paths may use for an op that requests an F32 src1, such as Mistral 4 `ffn_down_exps`. The default runs such ops on the library f32 GEMM; `TF32` or `BF16` trade mantissa for speed while keeping the f32 range. `F16` ignores the request and can overflow; it is meant for testing only. |
| GGML_SYCL_MMVQ_WIDE | 0 or 1 (default) | Use the wide-load variant of the reordered Q8_0 mat-vec kernel, which reads four contiguous dwords per operand instead of one value at a time. Set to 0 to fall back to the per-value loads. Only affects Q8_0 weights in the reordered layout. |
| GGML_SYCL_SPARSE_FA | 0 (default) or 1 | Enable Sparse Flash-attention.|
| GGML_SYCL_SPARSE_FA_DEBUG | 0 (default) or 1 | Enable to debug for Sparse Flash-attention.|
+2 -2
View File
@@ -164,11 +164,11 @@ export ZENDNNL_MATMUL_ALGO=1 # Blocked AOCL DLP algo for best performance
./build/bin/llama-server \
-m models/Llama-3.1-8B-Instruct.BF16.gguf \
--host 0.0.0.0 \
--port 8080 \
--port 9931 \
-t 64
```
Access the server at `http://localhost:8080`.
Access the server at `http://localhost:9931`.
**Performance tips**:
- Use `ZENDNNL_MATMUL_ALGO=1` for optimal performance
+1 -1
View File
@@ -351,7 +351,7 @@ cmake --build build --config Release
#### Override Compute Capability Specifications
By default, all supported compute capabilities are enabled. To customize this behavior, you can specify the `MUSA_ARCHITECTURES` option in the CMake command:
By default, compute capabilities `2.2` (MTT S4000) and `3.1` (MTT S5000) are enabled, compute capability `2.1` (MTT S70, MTT S80, MTT S3000) is deprecated and has to be enabled explicitly. To customize this behavior, you can specify the `MUSA_ARCHITECTURES` option in the CMake command:
```bash
cmake -B build -DGGML_MUSA=ON -DMUSA_ARCHITECTURES="31"
+4 -3
View File
@@ -25,15 +25,16 @@ output from a model that emits arguments as JSON.
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
// Build a choice of all available tools
auto tool_choice = p.choice();
for (const auto & tool : tools) {
for (size_t i = 0; i < tools.size(); i++) {
const auto & tool = tools[i];
const auto & function = tool.at("function");
std::string name = function.at("name");
const auto schema = common_chat_tool_parameters(function);
auto tool_name = p.json_member("name", "\"" + p.literal(name) + "\"");
auto tool_args = p.json_member("arguments", p.schema(p.json(), "tool-" + name + "-schema", schema));
auto tool_args = p.json_member("arguments", p.schema(p.json(), "tool-" + std::to_string(i) + "-schema", schema));
tool_choice |= p.rule("tool-" + name, "{" << tool_name << "," << tool_args << "}");
tool_choice |= p.rule("tool-" + std::to_string(i), "{" << tool_name << "," << tool_args << "}");
}
// Define the tool call structure: <tool_call>[{tool}]</tool_call>
+3 -3
View File
@@ -282,7 +282,7 @@ This table can be generated with:
# Usage - need tool-aware Jinja template
First, start a server with any model, but make sure it has a tools-enabled template: you can verify this by inspecting the `chat_template` or `chat_template_tool_use` properties in `http://localhost:8080/props`).
First, start a server with any model, but make sure it has a tools-enabled template: you can verify this by inspecting the `chat_template` or `chat_template_tool_use` properties in `http://localhost:9931/props`).
Here are some models known to work (w/ chat template override when needed):
@@ -336,7 +336,7 @@ To get the official template from original HuggingFace repos, you can use [scrip
Test in CLI (or with any library / software that can use OpenAI-compatible API backends):
```bash
curl http://localhost:8080/v1/chat/completions -d '{
curl http://localhost:9931/v1/chat/completions -d '{
"model": "gpt-3.5-turbo",
"tools": [
{
@@ -366,7 +366,7 @@ curl http://localhost:8080/v1/chat/completions -d '{
}'
curl http://localhost:8080/v1/chat/completions -d '{
curl http://localhost:9931/v1/chat/completions -d '{
"model": "gpt-3.5-turbo",
"messages": [
{"role": "system", "content": "You are a chatbot that uses tools/functions. Dont overthink things."},
+2 -2
View File
@@ -10,7 +10,7 @@ import json, requests
if True:
def create_completion(*, response_model=None, endpoint="http://localhost:8080/v1/chat/completions", messages, **kwargs):
def create_completion(*, response_model=None, endpoint="http://localhost:9931/v1/chat/completions", messages, **kwargs):
'''
Creates a chat completion using an OpenAI-compatible endpoint w/ JSON schema support
(llama.cpp server, llama-cpp-python, Anyscale / Together...)
@@ -45,7 +45,7 @@ else:
#! pip install instructor openai
import instructor, openai
client = instructor.patch(
openai.OpenAI(api_key="123", base_url="http://localhost:8080"),
openai.OpenAI(api_key="123", base_url="http://localhost:9931"),
mode=instructor.Mode.JSON_SCHEMA)
create_completion = client.chat.completions.create
@@ -10,4 +10,4 @@ Recommended way to run this model:
llama-server -hf {namespace}/{model_name}-GGUF
```
Then, access http://localhost:8080
Then, access http://localhost:9931
@@ -10,11 +10,11 @@ Recommended way to run this model:
llama-server -hf {namespace}/{model_name}-GGUF --embeddings
```
Then the endpoint can be accessed at http://localhost:8080/embedding, for
Then the endpoint can be accessed at http://localhost:9931/embedding, for
example using `curl`:
```console
curl --request POST \
--url http://localhost:8080/embedding \
--url http://localhost:9931/embedding \
--header "Content-Type: application/json" \
--data '{{"input": "Hello embeddings"}}' \
--silent
@@ -1,6 +1,6 @@
#!/usr/bin/env bash
curl --request POST \
--url http://localhost:8080/embedding \
--url http://localhost:9931/embedding \
--header "Content-Type: application/json" \
--data '{"input": "Hello world today"}' \
--silent
@@ -295,7 +295,7 @@ def example_concurrent(host):
def main():
parser = argparse.ArgumentParser(description=sys.modules[__name__].__doc__)
parser.add_argument("--host", default="localhost:8080", help="llama.cpp server")
parser.add_argument("--host", default="localhost:9931", help="llama.cpp server")
parser.add_argument("-v", "--verbose", action="store_true", help="enables logging")
args = parser.parse_args()
logging.basicConfig(level=logging.INFO if args.verbose else logging.ERROR)
@@ -206,7 +206,7 @@ int main(int argc, char ** argv) {
// reset the draft context to the checkpoint before verification
if (ctx_dft) {
if (use_ckpt_dft) {
ckpt.load_dft(ctx_dft, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
GGML_ASSERT(ckpt.load_dft(ctx_dft, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY));
}
llama_memory_seq_rm(llama_get_memory(ctx_dft), seq_id, ckpt.pos_max + 1, -1);
@@ -269,13 +269,13 @@ int main(int argc, char ** argv) {
draft = std::move(ids);
{
ckpt.load_tgt(ctx_tgt, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
GGML_ASSERT(ckpt.load_tgt(ctx_tgt, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY));
llama_memory_seq_rm(llama_get_memory(ctx_tgt), seq_id, ckpt.pos_max + 1, -1);
}
if (ctx_dft) {
ckpt.load_dft(ctx_dft, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
GGML_ASSERT(ckpt.load_dft(ctx_dft, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY));
llama_memory_seq_rm(llama_get_memory(ctx_dft), seq_id, ckpt.pos_max + 1, -1);
}
+9 -8
View File
@@ -2,7 +2,8 @@
#ifdef GGML_CUDA_USE_CUB
# include <cub/cub.cuh>
# if (CCCL_MAJOR_VERSION >= 3 && CCCL_MINOR_VERSION >= 1)
// strided_iterator was added in CCCL 3.1
# if (CCCL_MAJOR_VERSION > 3 || (CCCL_MAJOR_VERSION == 3 && CCCL_MINOR_VERSION >= 1))
# define STRIDED_ITERATOR_AVAILABLE
# include <cuda/iterator>
# endif
@@ -27,21 +28,21 @@ static __global__ void init_offsets(int * offsets, const int ncols, const int nr
}
#endif // STRIDED_ITERATOR_AVAILABLE
#ifdef GGML_CUDA_USE_CUB
// returns the suggested maximum number of rows to process during one argsort_f32_i32_cuda_cub() call
int argsort_f32_i32_cuda_cub_chunk_nrows(const size_t nb01, const int64_t nrows) {
// perform argsort in chunks up to approximately this size (currently 64MB)
// returns the suggested maximum number of rows to process at once, given the temporary buffer bytes per row
int ggml_cuda_chunk_nrows(const size_t row_bytes, const int64_t nrows) {
// process rows in chunks up to approximately this size (currently 64MB)
// to avoid excessive temporary buffers memory usage
const int chunk_bytes = 1 << 26;
// calculate how many rows will fit in one chunk (must be at least one)
const int chunk_nrows = std::max((int) (chunk_bytes / nb01), 1);
const int chunk_nrows = std::max((int) (chunk_bytes / row_bytes), 1);
// limit the resulting amount to total nrows
return std::min((int64_t) chunk_nrows, nrows);
}
#ifdef GGML_CUDA_USE_CUB
void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool,
const float * x,
int * dst,
@@ -289,7 +290,7 @@ void ggml_cuda_op_argsort(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
return;
}
const int chunk_nrows = argsort_f32_i32_cuda_cub_chunk_nrows(src0->nb[1], nrows);
const int chunk_nrows = ggml_cuda_chunk_nrows(src0->nb[1], nrows);
ggml_cuda_pool & pool = ctx.pool();
+2 -1
View File
@@ -4,8 +4,9 @@
void ggml_cuda_op_argsort(ggml_backend_cuda_context & ctx, ggml_tensor * dst);
int ggml_cuda_chunk_nrows(const size_t row_bytes, const int64_t nrows);
#ifdef GGML_CUDA_USE_CUB
int argsort_f32_i32_cuda_cub_chunk_nrows(const size_t nb01, const int64_t nrows);
void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool,
const float * x,
int * dst,
+115 -1
View File
@@ -2,6 +2,9 @@
#include "convert.cuh"
#include "fwht.cuh"
// wide FWHT blocks use one row per thread block with this many threads
#define GGML_CUDA_FWHT_BLOCK_NT 256
template <int N, typename T>
__launch_bounds__(4*ggml_cuda_get_physical_warp_size(), 1)
__global__ void fwht_cuda(const T * src, float * dst, const int64_t n_rows, const float scale) {
@@ -59,6 +62,87 @@ __global__ void fwht_cuda(const T * src, float * dst, const int64_t n_rows, cons
}
}
// Wide blocks: one row per thread block instead of per warp, so each thread keeps N/NT
// values rather than N/32. Stages below the warp width still shuffle, those up to the
// block width go through shared memory, and the rest stay in registers.
template <int N, int NT, typename T>
__launch_bounds__(NT, 1)
__global__ void fwht_cuda_block(const T * src, float * dst, const int64_t n_rows, const float scale) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int NE = N / NT;
static_assert(NE >= 1 && N % NT == 0 && NT % warp_size == 0, "bad FWHT block shape");
__shared__ float s[N];
const int64_t r = blockIdx.x;
if (r >= n_rows) {
return;
}
src += r * N;
dst += r * N;
const int tid = threadIdx.x;
const int lane = tid % warp_size;
ggml_cuda_pdl_sync();
float reg[NE];
#pragma unroll
for (int i = 0; i < NE; ++i) {
reg[i] = ggml_cuda_cast<float>(src[i * NT + tid]) * scale;
}
// stages within a warp: partner differs in the lane bits
#pragma unroll
for (int h = 1; h < warp_size; h *= 2) {
#pragma unroll
for (int j = 0; j < NE; j++) {
const float val = reg[j];
const float val2 = __shfl_xor_sync(0xFFFFFFFF, val, h, warp_size);
reg[j] = (lane & h) == 0 ? val + val2 : val2 - val;
}
}
// stages across warps: partner differs in the thread-index bits above the lane
#pragma unroll
for (int h = warp_size; h < NT; h *= 2) {
#pragma unroll
for (int j = 0; j < NE; j++) {
s[j * NT + tid] = reg[j];
}
__syncthreads();
#pragma unroll
for (int j = 0; j < NE; j++) {
const float val = reg[j];
const float val2 = s[j * NT + (tid ^ h)];
reg[j] = (tid & h) == 0 ? val + val2 : val2 - val;
}
__syncthreads();
}
// stages above the block width: partner is another register of the same thread
#pragma unroll
for (int h = NT; h < N; h *= 2) {
const int step = h / NT;
#pragma unroll
for (int j = 0; j < NE; j += 2 * step) {
#pragma unroll
for (int k = 0; k < step; k++) {
const float x = reg[j + k];
const float y = reg[j + k + step];
reg[j + k] = x + y;
reg[j + k + step] = x - y;
}
}
}
#pragma unroll
for (int i = 0; i < NE; ++i) {
dst[i * NT + tid] = reg[i];
}
}
template <typename T>
static bool ggml_cuda_op_fwht_impl(ggml_backend_cuda_context & ctx, const ggml_tensor * src, ggml_tensor * dst) {
const int n = src->ne[0];
@@ -94,7 +178,37 @@ static bool ggml_cuda_op_fwht_impl(ggml_backend_cuda_context & ctx, const ggml_t
ggml_cuda_kernel_launch(fwht_cuda<512, T>, launch_params, src_d, dst_d, rows, scale);
return true;
default:
return false;
break;
}
// wide blocks: one row per thread block
{
constexpr int nt = GGML_CUDA_FWHT_BLOCK_NT;
dim3 grid_dims_w(rows, 1, 1);
dim3 block_dims_w(nt, 1, 1);
const ggml_cuda_kernel_launch_params launch_params_w =
ggml_cuda_kernel_launch_params(grid_dims_w, block_dims_w, 0, stream);
switch (n) {
case 1024:
ggml_cuda_kernel_launch(fwht_cuda_block<1024, nt, T>, launch_params_w, src_d, dst_d, rows, scale);
return true;
case 2048:
ggml_cuda_kernel_launch(fwht_cuda_block<2048, nt, T>, launch_params_w, src_d, dst_d, rows, scale);
return true;
case 4096:
ggml_cuda_kernel_launch(fwht_cuda_block<4096, nt, T>, launch_params_w, src_d, dst_d, rows, scale);
return true;
#if !defined(GGML_USE_MUSA)
// 32 KB of shared memory, above the MUSA limit; falls back there
case 8192:
ggml_cuda_kernel_launch(fwht_cuda_block<8192, nt, T>, launch_params_w, src_d, dst_d, rows, scale);
return true;
#endif // !defined(GGML_USE_MUSA)
default:
return false;
}
}
}
+38 -21
View File
@@ -1,7 +1,9 @@
#include "gated_delta_net.cuh"
#include "ggml-cuda/common.cuh"
template <int S_v, bool KDA, bool keep_rs_t>
constexpr int gdn_cols_per_warp = 4;
template <int S_v, bool KDA, bool keep_rs_t, int cols_per_warp = gdn_cols_per_warp>
__global__ void __launch_bounds__((ggml_cuda_get_physical_warp_size() < S_v ? ggml_cuda_get_physical_warp_size() : S_v) * 4, 2)
gated_delta_net_cuda(const float * q,
const float * k,
@@ -30,9 +32,19 @@ gated_delta_net_cuda(const float * q,
int K) {
const uint32_t h_idx = blockIdx.x;
const uint32_t sequence = blockIdx.y;
// each warp owns one column, using warp-level primitives to reduce across rows
const int lane = threadIdx.x;
const int col = blockIdx.z * blockDim.y + threadIdx.y;
constexpr int warp_size = ggml_cuda_get_physical_warp_size() < S_v ? ggml_cuda_get_physical_warp_size() : S_v;
static_assert(S_v % warp_size == 0, "S_v must be a multiple of warp_size");
// the warp is split into cols_per_warp segments of lanes_per_col lanes; each segment owns
// one state column and reduces within itself
constexpr int lanes_per_col = warp_size / cols_per_warp;
constexpr int rows_per_lane = S_v / lanes_per_col;
static_assert(S_v % lanes_per_col == 0, "S_v must be a multiple of lanes_per_col");
const int lane = threadIdx.x;
const int col_in_warp = lane / lanes_per_col; // column slot within the warp
const int lane_in_col = lane - col_in_warp * lanes_per_col; // lane within the column's reduction segment
const int col = (blockIdx.z * blockDim.y + threadIdx.y) * cols_per_warp + col_in_warp;
const uint32_t iq1 = fastmodulo(h_idx, neqk1_magic);
const uint32_t iq3 = fastdiv(sequence, rq3_magic);
@@ -47,16 +59,13 @@ gated_delta_net_cuda(const float * q,
curr_state += state_in_offset + col * S_v;
attn_data += (sequence * n_tokens * H + h_idx) * S_v;
constexpr int warp_size = ggml_cuda_get_physical_warp_size() < S_v ? ggml_cuda_get_physical_warp_size() : S_v;
static_assert(S_v % warp_size == 0, "S_v must be a multiple of warp_size");
constexpr int rows_per_lane = (S_v + warp_size - 1) / warp_size;
float s_shard[rows_per_lane];
// state is stored transposed: M[col][i] = S[i][col], row col is contiguous
ggml_cuda_pdl_sync();
#pragma unroll
for (int r = 0; r < rows_per_lane; r++) {
const int i = r * warp_size + lane;
const int i = r * lanes_per_col + lane_in_col;
s_shard[r] = curr_state[i];
}
@@ -76,7 +85,7 @@ gated_delta_net_cuda(const float * q,
float q_reg[rows_per_lane];
#pragma unroll
for (int r = 0; r < rows_per_lane; r++) {
const int i = r * warp_size + lane;
const int i = r * lanes_per_col + lane_in_col;
k_reg[r] = k_t[i];
q_reg[r] = q_t[i];
}
@@ -90,7 +99,7 @@ gated_delta_net_cuda(const float * q,
for (int r = 0; r < rows_per_lane; r++) {
kv_shard += s_shard[r] * k_reg[r];
}
float kv_col = warp_reduce_sum<warp_size>(kv_shard);
float kv_col = warp_reduce_sum<lanes_per_col>(kv_shard);
// delta[col] = (v[col] - g * kv[col]) * beta
float delta_col = (v_t[col] - g_val * kv_col) * beta_val;
@@ -104,9 +113,9 @@ gated_delta_net_cuda(const float * q,
attn_partial += s_shard[r] * q_reg[r];
}
float attn_col = warp_reduce_sum<warp_size>(attn_partial);
float attn_col = warp_reduce_sum<lanes_per_col>(attn_partial);
if (lane == 0) {
if (lane_in_col == 0) {
attn_data[col] = attn_col * scale;
}
} else {
@@ -114,11 +123,11 @@ gated_delta_net_cuda(const float * q,
float kv_shard = 0.0f;
#pragma unroll
for (int r = 0; r < rows_per_lane; r++) {
const int i = r * warp_size + lane;
const int i = r * lanes_per_col + lane_in_col;
kv_shard += expf(g_t[i]) * s_shard[r] * k_reg[r];
}
float kv_col = warp_reduce_sum<warp_size>(kv_shard);
float kv_col = warp_reduce_sum<lanes_per_col>(kv_shard);
// delta[col] = (v[col] - kv[col]) * beta
float delta_col = (v_t[col] - kv_col) * beta_val;
@@ -128,14 +137,14 @@ gated_delta_net_cuda(const float * q,
float attn_partial = 0.0f;
#pragma unroll
for (int r = 0; r < rows_per_lane; r++) {
const int i = r * warp_size + lane;
const int i = r * lanes_per_col + lane_in_col;
s_shard[r] = expf(g_t[i]) * s_shard[r] + k_reg[r] * delta_col;
attn_partial += s_shard[r] * q_reg[r];
}
float attn_col = warp_reduce_sum<warp_size>(attn_partial);
float attn_col = warp_reduce_sum<lanes_per_col>(attn_partial);
if (lane == 0) {
if (lane_in_col == 0) {
attn_data[col] = attn_col * scale;
}
}
@@ -150,7 +159,7 @@ gated_delta_net_cuda(const float * q,
float * curr_state = state + target_slot * state_slot_stride;
#pragma unroll
for (int r = 0; r < rows_per_lane; r++) {
const int i = r * warp_size + lane;
const int i = r * lanes_per_col + lane_in_col;
curr_state[col * S_v + i] = s_shard[r];
}
}
@@ -160,7 +169,7 @@ gated_delta_net_cuda(const float * q,
if constexpr (!keep_rs_t) {
#pragma unroll
for (int r = 0; r < rows_per_lane; r++) {
const int i = r * warp_size + lane;
const int i = r * lanes_per_col + lane_in_col;
state[col * S_v + i] = s_shard[r];
}
}
@@ -179,8 +188,16 @@ static void launch_gated_delta_net(
float scale, int64_t state_slot_stride, int K, cudaStream_t stream) {
//TODO: Add chunked kernel for even faster pre-fill
const int warp_size = ggml_cuda_info().devices[ggml_cuda_get_device()].warp_size;
const int num_warps = 4;
dim3 grid_dims(H, n_seqs, (S_v + num_warps - 1) / num_warps);
// four columns per warp (see the kernel); shrink the CTA when the wider CTA would leave
// SMs without a CTA, so small head counts keep the device filled
const int nsm = ggml_cuda_info().devices[ggml_cuda_get_device()].nsm;
const int cols_per_warp = gdn_cols_per_warp;
int num_warps = 4;
while (num_warps > 1 && H*n_seqs*(S_v / (cols_per_warp * num_warps)) < nsm) {
num_warps /= 2;
}
// one CTA covers cols_per_warp*num_warps columns (see the kernel)
dim3 grid_dims(H, n_seqs, (S_v + cols_per_warp * num_warps - 1) / (cols_per_warp * num_warps));
dim3 block_dims(warp_size <= S_v ? warp_size : S_v, num_warps, 1);
const uint3 neqk1_magic = init_fastdiv_values(neqk1);
+7 -10
View File
@@ -5314,9 +5314,10 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
case GGML_UNARY_OP_CEIL:
case GGML_UNARY_OP_ROUND:
case GGML_UNARY_OP_TRUNC:
// TODO: should become:
//return ggml_is_contiguous_rows(op->src[0]);
return ggml_is_contiguous(op->src[0]);
if (op->src[0]->type == GGML_TYPE_BF16 && ggml_get_unary_op(op) == GGML_UNARY_OP_XIELU) {
return false;
}
return op->src[0]->type == GGML_TYPE_F16 || op->src[0]->type == GGML_TYPE_F32 || op->src[0]->type == GGML_TYPE_BF16;
default:
return false;
}
@@ -5658,7 +5659,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
return max_bias == 0.0f;
}
case GGML_OP_ROLL:
if(op->src[0]->type == GGML_TYPE_F32 && ggml_is_contiguous(op->src[0])) {
if(op->src[0]->type == GGML_TYPE_F32) {
return true;
}
return false;
@@ -5688,11 +5689,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
case GGML_OP_SUM:
return ggml_is_contiguous_rows(op->src[0]);
case GGML_OP_TOP_K:
#if defined(GGML_USE_HIP) || defined(GGML_CUDA_USE_CUB)
return true;
#else
return op->src[0]->ne[0] <= 1024;
#endif // defined(GGML_USE_HIP) || defined(GGML_CUDA_USE_CUB)
return op->src[0]->ne[0] <= INT_MAX;
case GGML_OP_ARGSORT:
#ifndef GGML_CUDA_USE_CUB
{
@@ -5704,7 +5701,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
return ncols_pad * sizeof(int) <= ggml_cuda_info().devices[dev_ctx->device].smpb;
}
#else
return true;
return op->src[0]->ne[0] <= INT_MAX;
#endif
case GGML_OP_SUM_ROWS:
return op->src[0]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32 && ggml_is_contiguous_rows(op->src[0]);
+64 -64
View File
@@ -7,9 +7,9 @@
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q1_0(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -98,9 +98,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q2_0(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -187,9 +187,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q4_0(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -250,9 +250,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q4_1(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -313,9 +313,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q5_0(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -393,9 +393,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q5_1(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -471,9 +471,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q8_0(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -537,9 +537,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q2_K(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -598,9 +598,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q3_K(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -711,9 +711,9 @@ static __device__ __forceinline__ int unpack_scales_q45_K(const int * scales, co
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q4_K(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -822,9 +822,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q5_K(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -946,9 +946,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q6_K(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1036,9 +1036,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq1_s(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1098,9 +1098,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq2_xxs(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1162,9 +1162,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq2_xs(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1227,9 +1227,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq2_s(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1295,9 +1295,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq3_xxs(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1359,9 +1359,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq3_s(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1428,9 +1428,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq4_xs(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1495,9 +1495,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq4_nl(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1564,9 +1564,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_mxfp4(
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1670,7 +1670,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
}
}
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_nvfp4(
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_nvfp4(
const char * __restrict__ x, int * __restrict__ x_tile, const int kb0, const int i_max, const int stride) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1) / warp_size;
+40 -40
View File
@@ -10,8 +10,8 @@ using namespace ggml_cuda_mma;
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q4_0_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q4_0, I);
const int * x_qs = (const int *) x;
@@ -60,8 +60,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q4_1_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q4_1, I);
const int * x_qs = (const int *) x;
@@ -110,8 +110,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q8_0, I);
const int * x_qs = (const int *) x;
@@ -148,8 +148,8 @@ static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma(
typedef tile<16, 8, int, input_layout> tile_B;
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -203,8 +203,8 @@ static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma(
typedef tile< 8, 8, int> tile_B;
typedef tile<16, 8, int> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -281,8 +281,8 @@ static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma(
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_1_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q5_1, I);
const int * x_qs = (const int *) x;
@@ -318,8 +318,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile<16, 8, int, input_layout> tile_B;
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -368,8 +368,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile< 8, 8, int> tile_B;
typedef tile<16, 8, int> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -442,8 +442,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(type, I);
const int * x_qs = (const int *) x;
@@ -474,7 +474,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
}
// Used for Q3_K, IQ2_S, and IQ2_XS:
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma(
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
#if defined(AMD_MFMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
constexpr data_layout input_layout = get_input_data_layout();
@@ -483,7 +483,7 @@ template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, prec_src1);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, prec_src1);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -533,7 +533,7 @@ template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_
typedef tile<16, 8, int> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, prec_src1);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, prec_src1);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -610,8 +610,8 @@ template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q2_K_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q2_K, I);
const int * x_qs = (const int *) x;
@@ -680,8 +680,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile<16, 4, int, input_layout> tile_B;
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -749,8 +749,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile< 8, 4, int> tile_B;
typedef tile<16, 8, int> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -870,8 +870,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q3_K_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q3_K, I);
const int * x_qs = (const int *) x;
@@ -905,8 +905,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q4_K_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q4_K, I);
const int * x_qs = (const int *) x;
@@ -940,8 +940,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q5_K_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q5_K, I);
const int * x_qs = (const int *) x;
@@ -975,8 +975,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q6_K_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q6_K, I);
const int * x_qs = (const int *) x;
@@ -1015,8 +1015,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile<16, 4, int, input_layout> tile_B;
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -1066,8 +1066,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile< 8, 4, int> tile_B;
typedef tile<16, 8, int> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -1181,7 +1181,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile<16, 8, float> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q4);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q4);
constexpr int ntx = rows_per_warp / tile_C::I;
constexpr int nfrags = MMQ_TILE_NE_K / tile_A::J;
+76 -38
View File
@@ -8,66 +8,66 @@
static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream, const ggml_prec prec_src1) {
switch (args.type_x) {
case GGML_TYPE_Q1_0:
mul_mat_q_case<GGML_TYPE_Q1_0>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_Q1_0, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q2_0:
mul_mat_q_case<GGML_TYPE_Q2_0>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_Q2_0, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q4_0:
mul_mat_q_case<GGML_TYPE_Q4_0>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_Q4_0, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q4_1:
mul_mat_q_case<GGML_TYPE_Q4_1>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_Q4_1, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q5_0:
mul_mat_q_case<GGML_TYPE_Q5_0>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_Q5_0, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q5_1:
mul_mat_q_case<GGML_TYPE_Q5_1>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_Q5_1, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q8_0:
mul_mat_q_case<GGML_TYPE_Q8_0>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_Q8_0, GGML_PREC_Q8>(ctx, args, stream);
break;
// -----------------------------------------------------------------------
case GGML_TYPE_Q2_K:
mul_mat_q_case<GGML_TYPE_Q2_K>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_Q2_K, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q3_K:
mul_mat_q_case<GGML_TYPE_Q3_K>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_Q3_K, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q4_K:
mul_mat_q_case<GGML_TYPE_Q4_K>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_Q4_K, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q5_K:
mul_mat_q_case<GGML_TYPE_Q5_K>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_Q5_K, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q6_K:
mul_mat_q_case<GGML_TYPE_Q6_K>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_Q6_K, GGML_PREC_Q8>(ctx, args, stream);
break;
// -----------------------------------------------------------------------
case GGML_TYPE_IQ1_S:
mul_mat_q_case<GGML_TYPE_IQ1_S>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_IQ1_S, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_IQ2_XXS:
mul_mat_q_case<GGML_TYPE_IQ2_XXS>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_IQ2_XXS, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_IQ2_XS:
mul_mat_q_case<GGML_TYPE_IQ2_XS>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_IQ2_XS, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_IQ2_S:
mul_mat_q_case<GGML_TYPE_IQ2_S>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_IQ2_S, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_IQ3_XXS:
mul_mat_q_case<GGML_TYPE_IQ3_XXS>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_IQ3_XXS, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_IQ3_S:
mul_mat_q_case<GGML_TYPE_IQ3_S>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_IQ3_S, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_IQ4_XS:
mul_mat_q_case<GGML_TYPE_IQ4_XS>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_IQ4_XS, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_IQ4_NL:
mul_mat_q_case<GGML_TYPE_IQ4_NL>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_IQ4_NL, GGML_PREC_Q8>(ctx, args, stream);
break;
// -----------------------------------------------------------------------
case GGML_TYPE_MXFP4:
@@ -76,14 +76,14 @@ static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, con
mul_mat_q_case<GGML_TYPE_MXFP4, GGML_PREC_Q4>(ctx, args, stream);
break;
}
mul_mat_q_case<GGML_TYPE_MXFP4>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_MXFP4, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_NVFP4:
if (prec_src1 == GGML_PREC_Q4) {
mul_mat_q_case<GGML_TYPE_NVFP4, GGML_PREC_Q4>(ctx, args, stream);
break;
}
mul_mat_q_case<GGML_TYPE_NVFP4>(ctx, args, stream);
mul_mat_q_case<GGML_TYPE_NVFP4, GGML_PREC_Q8>(ctx, args, stream);
break;
default:
GGML_ABORT("fatal error");
@@ -141,7 +141,10 @@ void ggml_cuda_mul_mat_q(
GGML_TENSOR_BINARY_OP_LOCALS;
cudaStream_t stream = ctx.stream();
const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc;
const int id = ggml_cuda_get_device();
const int cc = ggml_cuda_info().devices[id].cc;
const size_t smpbo = ggml_cuda_info().devices[id].smpbo;
const size_t ts_src0 = ggml_type_size(src0->type);
const size_t ts_src1 = ggml_type_size(src1->type);
@@ -176,7 +179,7 @@ void ggml_cuda_mul_mat_q(
const int64_t s03 = src0->nb[3] / ts_src0;
const int64_t s3 = dst->nb[3] / ts_dst;
const bool fallback = ne01 % 128 != 0;
const bool fallback = ggml_cuda_mmq_needs_fallback(ne01);
const ggml_prec prec_src1 = ggml_cuda_mmq_get_prec_src1(src0, dst, cc);
@@ -184,9 +187,52 @@ void ggml_cuda_mul_mat_q(
const size_t y_block_size = use_native_fp4 ? sizeof(block_fp4_mmq) : sizeof(block_q8_1_mmq);
const size_t y_values_per_block = use_native_fp4 ? QK_FP4_MMQ : QK8_1_MMQ;
int J_best = 0;
int nthreads_best = 0;
{
int64_t ncols_opt = ne11;
if (ids) {
const int64_t n_expert_used = ids->ne[0];
ncols_opt = ne12;
// Each expert only sees ne12*n_expert_used/ne02 tokens on average.
// On RDNA3 and RDNA4 it is faster to pick the tile size against this value instead of ne12.
if (GGML_CUDA_CC_IS_RDNA3(cc) || GGML_CUDA_CC_IS_RDNA4(cc)) {
ncols_opt = (ne12*n_expert_used + ne02 - 1) / ne02;
}
}
int ntiles_J_best = INT_MAX;
for (int J = 8; J <= 128 && ntiles_J_best > 1; J += 8) {
const ggml_cuda_mmq_config config = ggml_cuda_mmq_get_config(src0->type, J, fallback, cc, prec_src1);
if (config.type == GGML_TYPE_COUNT) {
continue;
}
if (mmq_get_nbytes_shared(config, cc) > smpbo) {
continue;
}
const int ntiles_x = (ncols_opt + config.J - 1) / config.J;
if (ntiles_x < ntiles_J_best) {
J_best = J;
nthreads_best = config.nthreads;
ntiles_J_best = ntiles_x;
}
}
}
GGML_ASSERT(J_best > 0);
// A tile of size J can read in at most J - 1 extra columns.
// For simplicity, round up the padding of a full tile to a multiple of the number of bytes that nthreads can load in parallel.
const size_t src1_load_chunk_size = nthreads_best * sizeof(int);
const size_t src1_q8_1_padding = ((J_best * sizeof(block_q8_1_mmq) + src1_load_chunk_size - 1) / src1_load_chunk_size)
* src1_load_chunk_size;
if (!ids) {
const size_t nbytes_src1_q8_1 = ne13*ne12 * ne11*ne10_padded * y_block_size/y_values_per_block +
ggml_cuda_mmq_get_J_max(src0->type, fallback, cc, ne11) * sizeof(block_q8_1_mmq);
const size_t nbytes_src1_q8_1 = ne13*ne12 * ne11*ne10_padded * y_block_size/y_values_per_block + src1_q8_1_padding;
ggml_cuda_pool_alloc<char> src1_q8_1(ctx.pool(), nbytes_src1_q8_1);
ggml_cuda_pool_alloc<float> src1_scale(ctx.pool());
if (src0->type == GGML_TYPE_NVFP4 && use_native_fp4) {
@@ -223,7 +269,7 @@ void ggml_cuda_mul_mat_q(
ne00, ne01, ne1, s01, ne11, s1,
ne02, ne12, s02, s12, s2,
ne03, ne13, s03, s13, s3,
ne1, ne1};
ne1, J_best};
ggml_cuda_mul_mat_q_switch_type(ctx, args, stream, prec_src1);
return;
}
@@ -237,7 +283,7 @@ void ggml_cuda_mul_mat_q(
GGML_ASSERT(ne1 == n_expert_used);
ggml_cuda_pool_alloc<int32_t> ids_src1(ctx.pool(), ne_get_rows);
ggml_cuda_pool_alloc<int32_t> ids_dst(ctx.pool(), ne_get_rows);
ggml_cuda_pool_alloc<int32_t> ids_dst(ctx.pool(), ne_get_rows + J_best-1); // Needs to be padded for unconditional memory access.
ggml_cuda_pool_alloc<int32_t> expert_bounds(ctx.pool(), ne02 + 1);
// gate/up activations are broadcast across experts (ne11 == 1): quantize each token once and
@@ -254,8 +300,7 @@ void ggml_cuda_mul_mat_q(
CUDA_CHECK(cudaGetLastError());
}
const size_t nbytes_src1_q8_1 = ne12*n_expert_used*ne10_padded * y_block_size/y_values_per_block +
ggml_cuda_mmq_get_J_max(src0->type, fallback, cc, ne12) * sizeof(block_q8_1_mmq);
const size_t nbytes_src1_q8_1 = ne12*n_expert_used*ne10_padded * y_block_size/y_values_per_block + src1_q8_1_padding;
ggml_cuda_pool_alloc<char> src1_q8_1(ctx.pool(), nbytes_src1_q8_1);
ggml_cuda_pool_alloc<float> src1_scale(ctx.pool());
if (src0->type == GGML_TYPE_NVFP4 && use_native_fp4) {
@@ -296,13 +341,6 @@ void ggml_cuda_mul_mat_q(
ne11 * ne10_padded * sizeof(block_q8_1) / (QK8_1 * sizeof(int));
const int64_t s13 = ne12*s12;
// Each expert only sees ne12*n_expert_used/ne02 tokens on average.
// On RDNA3 and RDNA4 it is faster to pick the tile size against this value instead of ne12.
int64_t ncols_opt = ne12;
if (GGML_CUDA_CC_IS_RDNA3(cc) || GGML_CUDA_CC_IS_RDNA4(cc)) {
ncols_opt = (ne12*n_expert_used + ne02 - 1) / ne02;
}
// Note that ne02 is used instead of ne12 because the number of y channels determines the z dimension of the CUDA grid.
const mmq_args args = {
src0_d, src0->type, (const int *) src1_q8_1.get(), ids_dst.get(), expert_bounds.get(), dst_d,
@@ -310,7 +348,7 @@ void ggml_cuda_mul_mat_q(
ne00, ne01, ne_get_rows, s01, ne_get_rows, s1,
ne02, ne02, s02, s12, s2,
ne03, ne13, s03, s13, s3,
ne12, ncols_opt};
ne12, J_best};
ggml_cuda_mul_mat_q_switch_type(ctx, args, stream, prec_src1);
}
+106 -138
View File
@@ -208,7 +208,7 @@ struct ggml_cuda_mmq_config {
static_assert((nthreads_) % 32 == 0 && (nthreads_) <= 512, "bad nthreads"); \
static_assert( (occupancy_) <= 8, "bad occupancy"); \
static_assert((I_) % 32 == 0, "bad I"); \
static_assert((J_) % 8 == 0, "bad J"); \
static_assert((J_) % 8 == 0 && (J_) <= 128, "bad J"); \
static_assert((K_vram_) % 256 == 0, "bad K_vram"); \
return ggml_cuda_mmq_config((type_), (nthreads_), (occupancy_), (I_), (J_), (sram_layout_), (K_vram_), (stream_k_), (fallback_)); \
} \
@@ -227,7 +227,7 @@ struct ggml_cuda_mmq_config {
#undef CASE
static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1 = GGML_PREC_Q8) {
static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
if (GGML_CUDA_CC_IS_AMD(cc)) {
if (GGML_CUDA_CC_IS_GCN(cc)) {
return ggml_cuda_mmq_get_config_gcn(type, J, fallback);
@@ -262,7 +262,7 @@ static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type ty
return ggml_cuda_mmq_get_config_pascal_older(type, J, fallback);
}
static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
#ifdef GGML_USE_HIP
#ifdef GCN
return ggml_cuda_mmq_get_config_gcn(type, J, fallback);
@@ -295,93 +295,86 @@ static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_t
GGML_UNUSED_VARS(type, J, fallback, prec_src1);
}
static __host__ int ggml_cuda_mmq_get_type(const ggml_type type, const int J, const bool fallback, const int cc) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc).type;
static __host__ int ggml_cuda_mmq_get_type(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).type;
}
static constexpr __device__ int ggml_cuda_mmq_get_type(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
static constexpr __device__ int ggml_cuda_mmq_get_type(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).type;
}
static constexpr __device__ int ggml_cuda_mmq_get_nthreads(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
static constexpr __device__ int ggml_cuda_mmq_get_nthreads(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).nthreads;
}
static constexpr __device__ int ggml_cuda_mmq_get_occupancy(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
static constexpr __device__ int ggml_cuda_mmq_get_occupancy(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).occupancy;
}
static __host__ int ggml_cuda_mmq_get_I(const ggml_type type, const int J, const bool fallback, const int cc) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc).I;
static __host__ int ggml_cuda_mmq_get_I(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).I;
}
static constexpr __device__ int ggml_cuda_mmq_get_I(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
static constexpr __device__ int ggml_cuda_mmq_get_I(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).I;
}
static __host__ int ggml_cuda_mmq_get_J(const ggml_type type, const int J, const bool fallback, const int cc) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc).J;
static __host__ int ggml_cuda_mmq_get_J(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).J;
}
static constexpr __device__ int ggml_cuda_mmq_get_J(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
static constexpr __device__ int ggml_cuda_mmq_get_J(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).J;
}
static __host__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(const ggml_type type, const int J, const bool fallback, const int cc) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc).sram_layout;
static __host__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).sram_layout;
}
static constexpr __device__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
static constexpr __device__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).sram_layout;
}
static __host__ int ggml_cuda_mmq_get_K_vram(const ggml_type type, const int J, const bool fallback, const int cc) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc).K_vram;
static __host__ int ggml_cuda_mmq_get_K_vram(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).K_vram;
}
static constexpr __device__ int ggml_cuda_mmq_get_K_vram(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
static constexpr __device__ int ggml_cuda_mmq_get_K_vram(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).K_vram;
}
static __host__ bool ggml_cuda_mmq_get_stream_k(const ggml_type type, const int J, const bool fallback, const int cc) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc).stream_k;
static __host__ bool ggml_cuda_mmq_get_stream_k(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).stream_k;
}
static constexpr __device__ bool ggml_cuda_mmq_get_stream_k(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
static constexpr __device__ bool ggml_cuda_mmq_get_stream_k(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).stream_k;
}
static __host__ int ggml_cuda_mmq_get_fallback(const ggml_type type, const int J, const bool fallback, const int cc) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc).fallback;
static __host__ int ggml_cuda_mmq_get_fallback(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).fallback;
}
static constexpr __device__ int ggml_cuda_mmq_get_fallback(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
static constexpr __device__ int ggml_cuda_mmq_get_fallback(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).fallback;
}
// ---------------------------------------------------------------------------------------------
static __host__ int ggml_cuda_mmq_get_sram_stride(const ggml_type type, const int J, const bool fallback, const int cc) {
return ggml_cuda_mmq_get_sram_stride(ggml_cuda_mmq_get_sram_layout(type, J, fallback, cc));
static __host__ int ggml_cuda_mmq_get_sram_stride(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
return ggml_cuda_mmq_get_sram_stride(ggml_cuda_mmq_get_sram_layout(type, J, fallback, cc, prec_src1));
}
static constexpr __device__ int ggml_cuda_mmq_get_sram_stride(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
static constexpr __device__ int ggml_cuda_mmq_get_sram_stride(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_sram_stride(ggml_cuda_mmq_get_sram_layout(type, J, fallback, prec_src1));
}
static __host__ int ggml_cuda_mmq_get_J_max(const ggml_type type, const bool fallback, const int cc, const int64_t ne11) {
int ret = std::min(ne11, int64_t(512));
ret -= ret % 8;
for (;ret > 0; ret -= 8) {
if (ggml_cuda_mmq_get_config(type, ret, fallback, cc).type != GGML_TYPE_COUNT) {
return ret;
}
}
return ret;
static __host__ bool ggml_cuda_mmq_needs_fallback(const int64_t nrows_x) {
return nrows_x % 128 != 0;
}
static constexpr __device__ int ggml_cuda_mmq_get_rows_per_warp(ggml_type type, int J, bool fallback) {
return ggml_cuda_mmq_get_config(type, J, fallback).rows_per_warp();
static constexpr __device__ int ggml_cuda_mmq_get_rows_per_warp(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).rows_per_warp();
}
#define MMQ_DP4A_TXS_Q4_0 tile_x_sizes{I*MMQ_TILE_NE_K + I, I*MMQ_TILE_NE_K/QI4_0 + I/QI4_0, 0}
@@ -437,12 +430,12 @@ static __host__ int ggml_cuda_mmq_get_nbytes_shared_x(const ggml_cuda_mmq_config
#include "mmq-load-tiles.cuh"
#include "mmq-vec-dot.cuh"
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_write_back_dp4a(
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1> static __device__ __forceinline__ void ggml_cuda_mmq_write_back_dp4a(
const float * __restrict__ sum, const int32_t * __restrict__ ids_dst, float * __restrict__ dst,
const float * __restrict__ y_scale, const int stride, const int i_max, const int j_max) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1) / warp_size;
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, prec_src1);
const bool y_scale_used = y_scale != nullptr;
@@ -476,7 +469,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
}
}
template<ggml_type type, int J, bool fallback>
template<ggml_type type, int J, bool fallback, ggml_prec prec_src1>
static __device__ __forceinline__ void ggml_cuda_mmq_write_back_mma(
const float * __restrict__ sum, const int * __restrict__ ids_dst, float * __restrict__ dst,
const float * __restrict__ y_scale, const int stride, const int i_max, const int j_max) {
@@ -487,7 +480,7 @@ static __device__ __forceinline__ void ggml_cuda_mmq_write_back_mma(
typedef tile<16, 8, int> tile_C;
#endif // defined(AMD_MFMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, prec_src1);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
const int i0 = (threadIdx.y / ntx) * (ntx*tile_C::I);
@@ -541,7 +534,7 @@ struct ggml_cuda_mmq_util_funcs {
vdr(vdr), load_tiles(load_tiles), vec_dot(vec_dot), write_back(write_back) {}
};
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_funcs() {
if (!ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).use_mma_data_layout()) {
switch (type) {
@@ -550,136 +543,136 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
VDR_Q1_0_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q1_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q2_0:
return ggml_cuda_mmq_util_funcs(
VDR_Q2_0_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q2_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q4_0:
return ggml_cuda_mmq_util_funcs(
VDR_Q4_0_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q4_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q4_0_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q4_1:
return ggml_cuda_mmq_util_funcs(
VDR_Q4_1_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q4_1<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q4_1_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q5_0:
return ggml_cuda_mmq_util_funcs(
VDR_Q5_0_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q5_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q5_1:
return ggml_cuda_mmq_util_funcs(
VDR_Q5_1_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q5_1<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_1_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q8_0:
return ggml_cuda_mmq_util_funcs(
VDR_Q8_0_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q8_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
// ---------------------------------------------------------------------------------------------
case GGML_TYPE_Q2_K:
return ggml_cuda_mmq_util_funcs(
VDR_Q2_K_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q2_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q2_K_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q3_K:
return ggml_cuda_mmq_util_funcs(
VDR_Q3_K_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q3_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q3_K_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q4_K:
return ggml_cuda_mmq_util_funcs(
VDR_Q4_K_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q4_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q4_K_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q5_K:
return ggml_cuda_mmq_util_funcs(
VDR_Q5_K_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q5_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q5_K_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q6_K:
return ggml_cuda_mmq_util_funcs(
VDR_Q6_K_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q6_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q6_K_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
// ---------------------------------------------------------------------------------------------
case GGML_TYPE_IQ1_S:
return ggml_cuda_mmq_util_funcs(
VDR_IQ1_S_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq1_s<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_1_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ2_XXS:
return ggml_cuda_mmq_util_funcs(
VDR_IQ2_XXS_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq2_xxs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ2_XS:
return ggml_cuda_mmq_util_funcs(
VDR_IQ2_XS_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq2_xs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ2_S:
return ggml_cuda_mmq_util_funcs(
VDR_IQ2_S_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq2_s<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ3_XXS:
return ggml_cuda_mmq_util_funcs(
VDR_IQ3_XXS_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq3_xxs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ3_S:
return ggml_cuda_mmq_util_funcs(
VDR_IQ3_S_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq3_s<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ4_XS:
return ggml_cuda_mmq_util_funcs(
VDR_IQ4_XS_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq4_xs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ4_NL:
return ggml_cuda_mmq_util_funcs(
VDR_IQ4_NL_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq4_nl<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
// ---------------------------------------------------------------------------------------------
case GGML_TYPE_MXFP4:
return ggml_cuda_mmq_util_funcs(
VDR_MXFP4_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_mxfp4<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_NVFP4:
return ggml_cuda_mmq_util_funcs(
VDR_NVFP4_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_nvfp4<type, J, fallback>,
ggml_cuda_mmq_load_tiles_nvfp4<type, J, fallback, prec_src1>,
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_dp4a<type, J, fallback>,
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
default:
return ggml_cuda_mmq_util_funcs(1, nullptr, nullptr, nullptr);
}
@@ -695,7 +688,7 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
-1,
ggml_cuda_mmq_load_tiles_mxfp4_fp4<type, J, fallback>,
ggml_cuda_mmq_vec_dot_fp4_fp4_mma<type, J, fallback>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
}
break;
case GGML_TYPE_NVFP4:
@@ -704,7 +697,7 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
-1,
ggml_cuda_mmq_load_tiles_nvfp4_nvfp4<type, J, fallback>,
ggml_cuda_mmq_vec_dot_fp4_fp4_mma<type, J, fallback>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
}
break;
default:
@@ -720,164 +713,164 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
-1,
ggml_cuda_mmq_load_tiles_q1_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q2_0:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q2_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q4_0:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q4_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_DS4>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q4_1:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q4_1<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q5_0:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q5_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q5_1:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q5_1<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q8_0:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q8_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
// ---------------------------------------------------------------------------------------------
case GGML_TYPE_Q2_K:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q2_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q2_K_q8_1_mma<type, J, fallback>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q3_K:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q3_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback, prec_src1>,
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q4_K:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q4_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q5_K:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q5_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q6_K:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q6_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q6_K_q8_1_mma<type, J, fallback>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
// ---------------------------------------------------------------------------------------------
case GGML_TYPE_IQ1_S:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq1_s<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ2_XXS:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq2_xxs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ2_XS:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq2_xs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback, prec_src1>,
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ2_S:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq2_s<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback, prec_src1>,
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ3_XXS:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq3_xxs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ3_S:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq3_s<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ4_XS:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq4_xs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ4_NL:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq4_nl<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
// ---------------------------------------------------------------------------------------------
case GGML_TYPE_MXFP4:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_mxfp4<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_NVFP4:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_nvfp4<type, J, fallback, prec_src1>,
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback, prec_src1>,
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
default:
return ggml_cuda_mmq_util_funcs(1, nullptr, nullptr, nullptr);
}
}
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
static constexpr __device__ int ggml_cuda_mmq_get_vdr() {
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().vdr;
}
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
static constexpr __device__ ggml_cuda_mmq_load_tiles_t ggml_cuda_mmq_get_load_tiles() {
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().load_tiles;
}
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
static constexpr __device__ ggml_cuda_mmq_vec_dot_t ggml_cuda_mmq_get_vec_dot() {
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().vec_dot;
}
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
static constexpr __device__ ggml_cuda_mmq_write_back_t ggml_cuda_mmq_get_write_back() {
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().write_back;
}
// ---------------------------------------------------------------------------------------------
template <ggml_type type, int J, bool fallback, bool fixup, ggml_prec prec_src1 = GGML_PREC_Q8>
template <ggml_type type, int J, bool fallback, bool fixup, ggml_prec prec_src1>
static __device__ __forceinline__ void mul_mat_q_process_tile(
const char * __restrict__ x, const int offset_x, const int * __restrict__ y,
const int * __restrict__ ids_dst, float * __restrict__ dst, float * __restrict__ tmp_fixup,
@@ -958,7 +951,7 @@ static __device__ __forceinline__ void mul_mat_q_process_tile(
// The mul_mat_q kernel implements "stream-k" work partitioning as described in https://arxiv.org/abs/2301.03598
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
__launch_bounds__(ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1), ggml_cuda_mmq_get_occupancy(type, J, fallback, prec_src1))
static __global__ void mul_mat_q(
const char * __restrict__ x, const int * __restrict__ y, const int32_t * __restrict__ ids_dst,
@@ -1245,7 +1238,7 @@ static __global__ void mul_mat_q(
tile_x_max_i, tile_y_max_j, kb0_start, kb0_stop);
}
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
__launch_bounds__(ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1)/2, 1)
static __global__ void mul_mat_q_stream_k_fixup(
const int32_t * __restrict__ ids_dst, const int32_t * __restrict__ expert_bounds, float * __restrict__ dst,
@@ -1390,7 +1383,7 @@ struct mmq_args {
int64_t nchannels_x; int64_t nchannels_y; int64_t stride_channel_x; int64_t stride_channel_y; int64_t stride_channel_dst;
int64_t nsamples_x; int64_t nsamples_y; int64_t stride_sample_x; int64_t stride_sample_y; int64_t stride_sample_dst;
int64_t ncols_max;
int64_t ncols_opt; // value to optimize the tile size against, launch grid still uses ncols_max
int J_best; // Tile width in ne11(dense)/ne12(MoE) direction to use for optimal performance.
};
static size_t mmq_get_nbytes_shared(const ggml_cuda_mmq_config & config, const int cc) {
@@ -1400,7 +1393,7 @@ static size_t mmq_get_nbytes_shared(const ggml_cuda_mmq_config & config, const i
return nbs_ids + nbs_x + GGML_PAD(nbs_y, config.nthreads*sizeof(int));
}
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
static void launch_mul_mat_q(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) {
const int id = ggml_cuda_get_device();
const int cc = ggml_cuda_info().devices[id].cc;
@@ -1482,34 +1475,9 @@ static void launch_mul_mat_q(ggml_backend_cuda_context & ctx, const mmq_args & a
ntx_fd);
}
template <ggml_type type, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
template <ggml_type type, bool fallback, ggml_prec prec_src1>
void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) {
const int id = ggml_cuda_get_device();
const int cc = ggml_cuda_info().devices[id].cc;
const size_t smpbo = ggml_cuda_info().devices[id].smpbo;
int J_best = 0;
int ntiles_J_best = INT_MAX;
for (int J = 8; J <= 128 && ntiles_J_best > 1; J += 8) {
const ggml_cuda_mmq_config config = ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1);
if (config.type == GGML_TYPE_COUNT) {
continue;
}
if (mmq_get_nbytes_shared(config, cc) > smpbo) {
continue;
}
const int ntiles_x = (args.ncols_opt + config.J - 1) / config.J;
if (ntiles_x < ntiles_J_best) {
J_best = J;
ntiles_J_best = ntiles_x;
}
}
switch (J_best) {
switch (args.J_best) {
case 8:
launch_mul_mat_q<type, 8, fallback, prec_src1>(ctx, args, stream);
break;
@@ -1559,25 +1527,25 @@ void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args,
launch_mul_mat_q<type, 128, fallback, prec_src1>(ctx, args, stream);
break;
default:
fprintf(stderr, "J_best=%d\n", J_best);
fprintf(stderr, "J_best=%d\n", args.J_best);
GGML_ABORT("fatal error");
break;
}
}
template <ggml_type type, ggml_prec prec_src1 = GGML_PREC_Q8>
template <ggml_type type, ggml_prec prec_src1>
void mul_mat_q_case(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) {
if (args.nrows_x % 128 == 0) {
constexpr bool fallback = false;
if (ggml_cuda_mmq_needs_fallback(args.nrows_x)) {
constexpr bool fallback = true;
mul_mat_q_switch_J<type, fallback, prec_src1>(ctx, args, stream);
} else {
constexpr bool fallback = true;
constexpr bool fallback = false;
mul_mat_q_switch_J<type, fallback, prec_src1>(ctx, args, stream);
}
}
#define DECL_MMQ_CASE(type) \
template void mul_mat_q_case<type>(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) \
template void mul_mat_q_case<type, GGML_PREC_Q8>(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) \
// FP4 variant: uses native FP4 MMA instead of keeping src1 at Q8_1.
#define DECL_MMQ_CASE_W4A4(type) \
+138 -111
View File
@@ -3,38 +3,46 @@
template <int block_size>
static __global__ void norm_f32(
const float * x, float * dst, const int ncols, const int64_t stride_row, const int64_t stride_channel,
const int64_t stride_sample, const float eps) {
const int nrows = gridDim.x;
const int nchannels = gridDim.y;
const float * x, float * dst, const int ncols, const int nchannels, const int nsamples,
const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps) {
const int nrows = gridDim.x;
const int row = blockIdx.x;
const int tid = threadIdx.x;
const int row = blockIdx.x;
const int channel = blockIdx.y;
const int sample = blockIdx.z;
const int tid = threadIdx.x;
x += sample*stride_sample + channel*stride_channel + row*stride_row;
dst += ((sample*nchannels + channel)*nrows + row)*ncols;
float2 mean_var = make_float2(0.0f, 0.0f);
extern __shared__ float2 s_sum2[];
ggml_cuda_pdl_sync();
for (int col = tid; col < ncols; col += block_size) {
const float xi = x[col];
mean_var.x += xi;
mean_var.y += xi * xi;
}
// sum up partial sums
extern __shared__ float2 s_sum2[];
mean_var = block_reduce<block_reduce_method::SUM, block_size>(mean_var, s_sum2);
// grid.y and grid.z are clamped to the CUDA limit, iterate over the excess channels/samples
for (int sample = blockIdx.z; sample < nsamples; sample += gridDim.z) {
for (int channel = blockIdx.y; channel < nchannels; channel += gridDim.y) {
const float * xc = x + sample*stride_sample + channel*stride_channel + row*stride_row;
float * dstc = dst + ((sample*nchannels + channel)*nrows + row)*ncols;
const float mean = mean_var.x / ncols;
const float var = mean_var.y / ncols - mean * mean;
const float inv_std = rsqrtf(var + eps);
float2 mean_var = make_float2(0.0f, 0.0f);
for (int col = tid; col < ncols; col += block_size) {
dst[col] = (x[col] - mean) * inv_std;
for (int col = tid; col < ncols; col += block_size) {
const float xi = xc[col];
mean_var.x += xi;
mean_var.y += xi * xi;
}
// sum up partial sums
mean_var = block_reduce<block_reduce_method::SUM, block_size>(mean_var, s_sum2);
const float mean = mean_var.x / ncols;
const float var = mean_var.y / ncols - mean * mean;
const float inv_std = rsqrtf(var + eps);
for (int col = tid; col < ncols; col += block_size) {
dstc[col] = (xc[col] - mean) * inv_std;
}
if constexpr (block_size > WARP_SIZE) {
// sync is needed as we reuse s_sum2 across block_reduce invocations, see #26385
__syncthreads();
}
}
}
}
@@ -77,6 +85,8 @@ template <int block_size, bool do_multiply = false, bool do_add = false, bool do
static __global__ void rms_norm_f32(const float * x,
float * dst,
const int ncols,
const int nchannels,
const int nsamples,
const int64_t stride_row,
const int64_t stride_channel,
const int64_t stride_sample,
@@ -99,61 +109,71 @@ static __global__ void rms_norm_f32(const float * x,
const uint3 add_nsamples_packed = make_uint3(0, 0, 0),
const float scale_out = 1.0f) {
ggml_cuda_pdl_lc();
const int nrows = gridDim.x;
const int nchannels = gridDim.y;
const int row = blockIdx.x;
const int channel = blockIdx.y;
const int sample = blockIdx.z;
const int tid = threadIdx.x;
const int nrows = gridDim.x;
const int row = blockIdx.x;
const int tid = threadIdx.x;
static_assert(!do_add || do_multiply, "fusing add is not supported without multiplying");
static_assert(!do_scale || !do_multiply, "fusing scale is not supported with multiplying");
x += sample*stride_sample + channel*stride_channel + row*stride_row;
dst += ((sample*nchannels + channel)*nrows + row)*ncols;
if constexpr (do_multiply) {
const uint32_t mul_row = fastmodulo(row, mul_nrows_packed);
const uint32_t mul_channel = fastmodulo(channel, mul_nchannels_packed);
const uint32_t mul_sample = fastmodulo(sample, mul_nsamples_packed);
mul += mul_sample * mul_stride_sample + mul_channel * mul_stride_channel + mul_row * mul_stride_row;
}
if constexpr (do_add) {
const int add_row = fastmodulo(row, add_nrows_packed);
const int add_channel = fastmodulo(channel, add_nchannels_packed);
const int add_sample = fastmodulo(sample, add_nsamples_packed);
add += add_sample * add_stride_sample + add_channel * add_stride_channel + add_row * add_stride_row;
}
float tmp = 0.0f; // partial sum for thread in warp
extern __shared__ float s_sum[];
ggml_cuda_pdl_sync();
for (int col = tid; col < ncols; col += block_size) {
const float xi = x[col];
tmp += xi * xi;
}
// sum up partial sums
extern __shared__ float s_sum[];
tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
// grid.y and grid.z are clamped to the CUDA limit, iterate over the excess channels/samples
for (int sample = blockIdx.z; sample < nsamples; sample += gridDim.z) {
for (int channel = blockIdx.y; channel < nchannels; channel += gridDim.y) {
const float * xc = x + sample*stride_sample + channel*stride_channel + row*stride_row;
float * dstc = dst + ((sample*nchannels + channel)*nrows + row)*ncols;
const float mean = tmp / ncols;
const float scale = rsqrtf(mean + eps);
[[maybe_unused]] const float * mulc = nullptr;
if constexpr (do_multiply) {
const uint32_t mul_row = fastmodulo(row, mul_nrows_packed);
const uint32_t mul_channel = fastmodulo(channel, mul_nchannels_packed);
const uint32_t mul_sample = fastmodulo(sample, mul_nsamples_packed);
mulc = mul + mul_sample * mul_stride_sample + mul_channel * mul_stride_channel + mul_row * mul_stride_row;
}
for (int col = tid; col < ncols; col += block_size) {
if constexpr (do_multiply && do_add) {
const int mul_col = fastmodulo(col, mul_ncols_packed);
const int add_col = fastmodulo(col, add_ncols_packed);
dst[col] = scale * x[col] * mul[mul_col] + add[add_col];
} else if constexpr (do_multiply) {
const int mul_col = fastmodulo(col, mul_ncols_packed);
dst[col] = scale * x[col] * mul[mul_col];
} else if constexpr (do_scale) {
dst[col] = scale_out * (scale * x[col]);
} else {
dst[col] = scale * x[col];
[[maybe_unused]] const float * addc = nullptr;
if constexpr (do_add) {
const int add_row = fastmodulo(row, add_nrows_packed);
const int add_channel = fastmodulo(channel, add_nchannels_packed);
const int add_sample = fastmodulo(sample, add_nsamples_packed);
addc = add + add_sample * add_stride_sample + add_channel * add_stride_channel + add_row * add_stride_row;
}
float tmp = 0.0f; // partial sum for thread in warp
for (int col = tid; col < ncols; col += block_size) {
const float xi = xc[col];
tmp += xi * xi;
}
// sum up partial sums
tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
const float mean = tmp / ncols;
const float scale = rsqrtf(mean + eps);
for (int col = tid; col < ncols; col += block_size) {
if constexpr (do_multiply && do_add) {
const int mul_col = fastmodulo(col, mul_ncols_packed);
const int add_col = fastmodulo(col, add_ncols_packed);
dstc[col] = scale * xc[col] * mulc[mul_col] + addc[add_col];
} else if constexpr (do_multiply) {
const int mul_col = fastmodulo(col, mul_ncols_packed);
dstc[col] = scale * xc[col] * mulc[mul_col];
} else if constexpr (do_scale) {
dstc[col] = scale_out * (scale * xc[col]);
} else {
dstc[col] = scale * xc[col];
}
}
if constexpr (block_size > WARP_SIZE) {
// sync is needed as we reuse s_sum across block_reduce invocations, see #26385
__syncthreads();
}
}
}
}
@@ -247,50 +267,57 @@ static __global__ void rms_norm_back_f32(
template <int block_size>
static __global__ void l2_norm_f32(
const float * x, float * dst, const int ncols, const int64_t stride_row, const int64_t stride_channel,
const int64_t stride_sample, const float eps) {
const int nrows = gridDim.x;
const int nchannels = gridDim.y;
const float * x, float * dst, const int ncols, const int nchannels, const int nsamples,
const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps) {
const int nrows = gridDim.x;
const int row = blockIdx.x;
const int tid = threadIdx.x;
const int row = blockIdx.x;
const int channel = blockIdx.y;
const int sample = blockIdx.z;
const int tid = threadIdx.x;
x += sample*stride_sample + channel*stride_channel + row*stride_row;
dst += ((sample*nchannels + channel)*nrows + row)*ncols;
float tmp = 0.0f; // partial sum for thread in warp
extern __shared__ float s_sum[];
ggml_cuda_pdl_sync();
for (int col = tid; col < ncols; col += block_size) {
const float xi = x[col];
tmp += xi * xi;
}
// sum up partial sums
extern __shared__ float s_sum[];
tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
ggml_cuda_pdl_lc();
// grid.y and grid.z are clamped to the CUDA limit, iterate over the excess channels/samples
for (int sample = blockIdx.z; sample < nsamples; sample += gridDim.z) {
for (int channel = blockIdx.y; channel < nchannels; channel += gridDim.y) {
const float * xc = x + sample*stride_sample + channel*stride_channel + row*stride_row;
float * dstc = dst + ((sample*nchannels + channel)*nrows + row)*ncols;
// from https://pytorch.org/docs/stable/generated/torch.nn.functional.normalize.html
const float scale = rsqrtf(fmaxf(tmp, eps * eps));
float tmp = 0.0f; // partial sum for thread in warp
for (int col = tid; col < ncols; col += block_size) {
dst[col] = scale * x[col];
for (int col = tid; col < ncols; col += block_size) {
const float xi = xc[col];
tmp += xi * xi;
}
// sum up partial sums
tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
// from https://pytorch.org/docs/stable/generated/torch.nn.functional.normalize.html
const float scale = rsqrtf(fmaxf(tmp, eps * eps));
for (int col = tid; col < ncols; col += block_size) {
dstc[col] = scale * xc[col];
}
if constexpr (block_size > WARP_SIZE) {
// sync is needed as we reuse s_sum across block_reduce invocations, see #26385
__syncthreads();
}
}
}
}
static void norm_f32_cuda(
const float * x, float * dst, const int ncols, const int nrows, const int nchannels, const int nsamples,
const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps, cudaStream_t stream) {
const dim3 blocks_num(nrows, nchannels, nsamples);
const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));
if (ncols < 1024) {
const dim3 block_dims(WARP_SIZE, 1, 1);
norm_f32<WARP_SIZE><<<blocks_num, block_dims, 0, stream>>>(x, dst, ncols, stride_row, stride_channel, stride_sample, eps);
norm_f32<WARP_SIZE><<<blocks_num, block_dims, 0, stream>>>(x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps);
} else {
const dim3 block_dims(1024, 1, 1);
norm_f32<1024><<<blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float2): 0, stream>>>(x, dst, ncols, stride_row, stride_channel, stride_sample, eps);
norm_f32<1024><<<blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float2): 0, stream>>>(x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps);
}
}
@@ -310,19 +337,19 @@ static void rms_norm_f32_cuda(
const float * x, float * dst, const int ncols, const int nrows, const int nchannels, const int nsamples,
const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps, cudaStream_t stream,
const float scale_out = 1.0f) {
const dim3 blocks_num(nrows, nchannels, nsamples);
const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));
if (ncols < 1024) {
const dim3 block_dims(256, 1, 1);
const ggml_cuda_kernel_launch_params launch_params = {blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
ggml_cuda_kernel_launch(rms_norm_f32<256, false, false, do_scale>, launch_params,
x, dst, ncols, stride_row, stride_channel, stride_sample, eps,
x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps,
// underlying cudaLaunchKernelEx does not support default params
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0),
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), scale_out);
} else {
const dim3 block_dims(1024, 1, 1);
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
ggml_cuda_kernel_launch(rms_norm_f32<1024, false, false, do_scale>, launch_params, x, dst, ncols, stride_row, stride_channel, stride_sample, eps,
ggml_cuda_kernel_launch(rms_norm_f32<1024, false, false, do_scale>, launch_params, x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps,
// underlying cudaLaunchKernelEx does not support default params
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0),
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), scale_out);
@@ -356,7 +383,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
const uint32_t add_nsamples,
const float eps,
cudaStream_t stream) {
const dim3 blocks_num(nrows, nchannels, nsamples);
const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));
if (mul == nullptr) {
rms_norm_f32_cuda(x, dst, ncols, nrows, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, stream);
return;
@@ -370,7 +397,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
const dim3 block_dims(256, 1, 1);
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
ggml_cuda_kernel_launch(rms_norm_f32<256, true>, launch_params,
x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
// underlying cudaLaunchKernelEx does not support default params
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), 1.0f);
@@ -378,7 +405,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
const dim3 block_dims(1024, 1, 1);
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
ggml_cuda_kernel_launch(rms_norm_f32<1024, true>, launch_params,
x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
// underlying cudaLaunchKernelEx does not support default params
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), 1.0f);
@@ -397,7 +424,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
const dim3 block_dims(256, 1, 1);
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims,block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
ggml_cuda_kernel_launch(rms_norm_f32<256, true, true>, launch_params,
x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed, add,
add_stride_row, add_stride_channel, add_stride_sample, add_ncols_packed, add_nrows_packed,
add_nchannels_packed, add_nsamples_packed, 1.0f);
@@ -405,7 +432,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
const dim3 block_dims(1024, 1, 1);
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
ggml_cuda_kernel_launch(rms_norm_f32<1024, true, true>, launch_params,
x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed, add,
add_stride_row, add_stride_channel, add_stride_sample, add_ncols_packed, add_nrows_packed,
add_nchannels_packed, add_nsamples_packed, 1.0f);
@@ -426,15 +453,15 @@ static void rms_norm_back_f32_cuda(const float * grad, const float * xf, float *
static void l2_norm_f32_cuda(
const float * x, float * dst, const int ncols, const int nrows, const int nchannels, const int nsamples,
const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps, cudaStream_t stream) {
const dim3 blocks_num(nrows, nchannels, nsamples);
const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));
if (ncols < 1024) {
const dim3 block_dims(WARP_SIZE, 1, 1);
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, 0, stream};
ggml_cuda_kernel_launch(l2_norm_f32<WARP_SIZE>, launch_params, x, dst, ncols, stride_row, stride_channel, stride_sample, eps);
ggml_cuda_kernel_launch(l2_norm_f32<WARP_SIZE>, launch_params, x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps);
} else {
const dim3 block_dims(1024, 1, 1);
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
ggml_cuda_kernel_launch(l2_norm_f32<1024>, launch_params, x, dst, ncols, stride_row, stride_channel, stride_sample, eps);
ggml_cuda_kernel_launch(l2_norm_f32<1024>, launch_params, x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps);
}
}
+38 -35
View File
@@ -15,49 +15,52 @@ static __global__ void pad_f32(const float * src, size_t s00, size_t s01, size_t
// blockIdx.z: i3*ne2+i2
// blockIdx.y: i1
// blockIDx.x: i0 / CUDA_PAD_BLOCK_SIZE
// gridDim.y: ne1
// gridDim.y and gridDim.z are capped at 65535, blocks stride over larger ne1 and ne2*ne3
int i0 = threadIdx.x + blockIdx.x * blockDim.x;
int i1 = blockIdx.y;
int i2 = blockIdx.z % ne2;
int i3 = blockIdx.z / ne2;
if (i0 >= ne0 || i1 >= ne1 || i2 >= ne2 || i3 >= ne3) {
if (i0 >= ne0) {
return;
}
const int64_t dst_idx = i3 * (ne0 * ne1 * ne2) + i2 * (ne0 * ne1) + i1 * ne0 + i0;
for (int i1 = blockIdx.y; i1 < ne1; i1 += gridDim.y) {
for (int i23 = blockIdx.z; i23 < ne2 * ne3; i23 += gridDim.z) {
int i2 = i23 % ne2;
int i3 = i23 / ne2;
if (!circular) {
if ((i0 >= lp0 && i0 < ne0 - rp0) && (i1 >= lp1 && i1 < ne1 - rp1) && (i2 >= lp2 && i2 < ne2 - rp2) &&
(i3 >= lp3 && i3 < ne3 - rp3)) {
const int64_t i00 = i0 - lp0;
const int64_t i01 = i1 - lp1;
const int64_t i02 = i2 - lp2;
const int64_t i03 = i3 - lp3;
const int64_t dst_idx = i3 * (ne0 * ne1 * ne2) + i2 * (ne0 * ne1) + i1 * ne0 + i0;
const int64_t src_idx = i03 * s03 + i02 * s02 + i01 * s01 + i00 * s00;
if (!circular) {
if ((i0 >= lp0 && i0 < ne0 - rp0) && (i1 >= lp1 && i1 < ne1 - rp1) && (i2 >= lp2 && i2 < ne2 - rp2) &&
(i3 >= lp3 && i3 < ne3 - rp3)) {
const int64_t i00 = i0 - lp0;
const int64_t i01 = i1 - lp1;
const int64_t i02 = i2 - lp2;
const int64_t i03 = i3 - lp3;
dst[dst_idx] = src[src_idx];
} else {
dst[dst_idx] = 0.0f;
const int64_t src_idx = i03 * s03 + i02 * s02 + i01 * s01 + i00 * s00;
dst[dst_idx] = src[src_idx];
} else {
dst[dst_idx] = 0.0f;
}
}
// circular means on a torus, so x and y wrap around
else {
const int64_t ne00 = ne0 - lp0 - rp0;
const int64_t ne01 = ne1 - lp1 - rp1;
const int64_t ne02 = ne2 - lp2 - rp2;
const int64_t ne03 = ne3 - lp3 - rp3;
const int64_t i00 = wrap_around(i0 - lp0, ne00);
const int64_t i01 = wrap_around(i1 - lp1, ne01);
const int64_t i02 = wrap_around(i2 - lp2, ne02);
const int64_t i03 = wrap_around(i3 - lp3, ne03);
const int64_t src_idx = i03 * s03 + i02 * s02 + i01 * s01 + i00 * s00;
dst[dst_idx] = src[src_idx];
}
}
}
// circular means on a torus, so x and y wrap around
else {
const int64_t ne00 = ne0 - lp0 - rp0;
const int64_t ne01 = ne1 - lp1 - rp1;
const int64_t ne02 = ne2 - lp2 - rp2;
const int64_t ne03 = ne3 - lp3 - rp3;
const int64_t i00 = wrap_around(i0 - lp0, ne00);
const int64_t i01 = wrap_around(i1 - lp1, ne01);
const int64_t i02 = wrap_around(i2 - lp2, ne02);
const int64_t i03 = wrap_around(i3 - lp3, ne03);
const int64_t src_idx = i03 * s03 + i02 * s02 + i01 * s01 + i00 * s00;
dst[dst_idx] = src[src_idx];
}
}
@@ -67,7 +70,7 @@ static void pad_f32_cuda(const float * src, size_t s00, size_t s01, size_t s02,
const int ne0, const int ne1, const int ne2, const int ne3,
const bool circular, cudaStream_t stream) {
int num_blocks = (ne0 + CUDA_PAD_BLOCK_SIZE - 1) / CUDA_PAD_BLOCK_SIZE;
dim3 gridDim(num_blocks, ne1, ne2 * ne3);
dim3 gridDim(num_blocks, std::min(ne1, 65535), std::min(ne2 * ne3, 65535));
pad_f32<<<gridDim, CUDA_PAD_BLOCK_SIZE, 0, stream>>>(src, s00, s01, s02, s03, dst,
lp0, rp0, lp1, rp1, lp2, rp2, lp3, rp3,
ne0, ne1, ne2, ne3, circular);
+6 -2
View File
@@ -17,6 +17,10 @@ static __global__ void roll_f32_cuda(const float * __restrict__ src,
const int64_t ne01,
const int64_t ne02,
const int64_t ne03,
const int64_t nb00,
const int64_t nb01,
const int64_t nb02,
const int64_t nb03,
const int s0,
const int s1,
const int s2,
@@ -39,7 +43,7 @@ static __global__ void roll_f32_cuda(const float * __restrict__ src,
const int64_t d3 = wrap_index(i3 - s3, ne03);
dst[i3 * (ne00 * ne01 * ne02) + i2 * (ne01 * ne00) + i1 * ne00 + i0] =
src[d3 * (ne00 * ne01 * ne02) + d2 * (ne01 * ne00) + d1 * ne00 + d0];
src[(d3 * nb03 + d2 * nb02 + d1 * nb01 + d0 * nb00) / sizeof(float)];
}
void ggml_cuda_op_roll(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
@@ -63,5 +67,5 @@ void ggml_cuda_op_roll(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
int64_t num_blocks = (sz + CUDA_ROLL_BLOCK_SIZE - 1) / CUDA_ROLL_BLOCK_SIZE;
roll_f32_cuda<<<num_blocks, CUDA_ROLL_BLOCK_SIZE, 0, stream>>>(
src0_d, dst_d, ne00, ne01, ne02, ne03, s0, s1, s2, s3);
src0_d, dst_d, ne00, ne01, ne02, ne03, nb00, nb01, nb02, nb03, s0, s1, s2, s3);
}
+70 -60
View File
@@ -709,7 +709,7 @@ void ggml_cuda_op_rope_fused(ggml_backend_cuda_context & ctx, ggml_tensor * rope
// one block per row: block_reduce gives the norm scale, then each thread applies mul and rope to the elements it owns
template <int block_size, bool has_ff, typename D>
static __global__ void rms_norm_mul_rope_f32(
const float * x, D * dst, const int ncols,
const float * x, D * dst, const int ncols, const int nchannels, const int nsamples,
const int64_t s01, const int64_t s02, const int64_t s03,
const int64_t s1, const int64_t s2, const int64_t s3,
const float eps,
@@ -724,66 +724,76 @@ static __global__ void rms_norm_mul_rope_f32(
const int64_t * row_indices, const int set_rows_stride,
const bool is_neox) {
ggml_cuda_pdl_lc();
const int row = blockIdx.x;
const int channel = blockIdx.y;
const int sample = blockIdx.z;
const int tid = threadIdx.x;
x += sample*s03 + channel*s02 + row*s01;
const uint32_t mul_row = fastmodulo(row, mul_nrows_packed);
const uint32_t mul_channel = fastmodulo(channel, mul_nchannels_packed);
const uint32_t mul_sample = fastmodulo(sample, mul_nsamples_packed);
mul += mul_sample*mul_s03 + mul_channel*mul_s02 + mul_row*mul_s01;
float tmp = 0.0f;
ggml_cuda_pdl_sync();
for (int col = tid; col < ncols; col += block_size) {
const float xi = x[col];
tmp += xi * xi;
}
const int row = blockIdx.x;
const int tid = threadIdx.x;
extern __shared__ float s_sum[];
tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
const float scale = rsqrtf(tmp/ncols + eps);
ggml_cuda_pdl_sync();
int64_t idst = sample*s3 + channel*s2 + row*s1;
if (set_rows_stride != 0) {
idst = row*s1 + row_indices[channel]*set_rows_stride;
}
dst += idst;
// grid.y and grid.z are clamped to the CUDA limit, iterate over the excess channels/samples
for (int sample = blockIdx.z; sample < nsamples; sample += gridDim.z) {
for (int channel = blockIdx.y; channel < nchannels; channel += gridDim.y) {
const float * xc = x + sample*s03 + channel*s02 + row*s01;
for (int i0 = 2*tid; i0 < ncols; i0 += 2*block_size) {
int ix0;
int ix1;
if (is_neox && i0 < n_dims) {
ix0 = i0/2;
ix1 = i0/2 + n_dims/2;
} else {
ix0 = i0 + 0;
ix1 = i0 + 1;
const uint32_t mul_row = fastmodulo(row, mul_nrows_packed);
const uint32_t mul_channel = fastmodulo(channel, mul_nchannels_packed);
const uint32_t mul_sample = fastmodulo(sample, mul_nsamples_packed);
const float * mulc = mul + mul_sample*mul_s03 + mul_channel*mul_s02 + mul_row*mul_s01;
float tmp = 0.0f;
for (int col = tid; col < ncols; col += block_size) {
const float xi = xc[col];
tmp += xi * xi;
}
tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
const float scale = rsqrtf(tmp/ncols + eps);
int64_t idst = sample*s3 + channel*s2 + row*s1;
if (set_rows_stride != 0) {
idst = row*s1 + row_indices[channel]*set_rows_stride;
}
D * dstc = dst + idst;
for (int i0 = 2*tid; i0 < ncols; i0 += 2*block_size) {
int ix0;
int ix1;
if (is_neox && i0 < n_dims) {
ix0 = i0/2;
ix1 = i0/2 + n_dims/2;
} else {
ix0 = i0 + 0;
ix1 = i0 + 1;
}
const float x0 = scale * xc[ix0] * mulc[fastmodulo(ix0, mul_ncols_packed)];
const float x1 = scale * xc[ix1] * mulc[fastmodulo(ix1, mul_ncols_packed)];
if (i0 >= n_dims) {
dstc[ix0] = ggml_cuda_cast<D>(x0);
dstc[ix1] = ggml_cuda_cast<D>(x1);
continue;
}
const float theta_base = pos[channel]*powf(theta_scale, i0/2.0f);
const float freq_factor = has_ff ? freq_factors[i0/2] : 1.0f;
float cos_theta;
float sin_theta;
rope_yarn<true>(theta_base/freq_factor, freq_scale, corr_dims, i0, ext_factor, attn_factor, cos_theta, sin_theta);
dstc[ix0] = ggml_cuda_cast<D>(x0*cos_theta - x1*sin_theta);
dstc[ix1] = ggml_cuda_cast<D>(x0*sin_theta + x1*cos_theta);
}
if constexpr (block_size > WARP_SIZE) {
// sync is needed as we reuse s_sum across block_reduce invocations, see #26385
__syncthreads();
}
}
const float x0 = scale * x[ix0] * mul[fastmodulo(ix0, mul_ncols_packed)];
const float x1 = scale * x[ix1] * mul[fastmodulo(ix1, mul_ncols_packed)];
if (i0 >= n_dims) {
dst[ix0] = ggml_cuda_cast<D>(x0);
dst[ix1] = ggml_cuda_cast<D>(x1);
continue;
}
const float theta_base = pos[channel]*powf(theta_scale, i0/2.0f);
const float freq_factor = has_ff ? freq_factors[i0/2] : 1.0f;
float cos_theta;
float sin_theta;
rope_yarn<true>(theta_base/freq_factor, freq_scale, corr_dims, i0, ext_factor, attn_factor, cos_theta, sin_theta);
dst[ix0] = ggml_cuda_cast<D>(x0*cos_theta - x1*sin_theta);
dst[ix1] = ggml_cuda_cast<D>(x0*sin_theta + x1*cos_theta);
}
}
@@ -806,7 +816,7 @@ static void rms_norm_mul_rope_cuda(
const bool is_neox, cudaStream_t stream) {
GGML_ASSERT(ncols % 2 == 0);
const dim3 blocks_num(nrows, nchannels, nsamples);
const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));
const float theta_scale = powf(freq_base, -2.0f/n_dims);
@@ -820,13 +830,13 @@ static void rms_norm_mul_rope_cuda(
const ggml_cuda_kernel_launch_params launch_params = {blocks_num, block_dims, 32*sizeof(float), stream};
if (freq_factors == nullptr) {
ggml_cuda_kernel_launch(rms_norm_mul_rope_f32<256, false, D>, launch_params,
x, dst, ncols, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
x, dst, ncols, nchannels, nsamples, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
n_dims, pos, freq_scale, ext_factor, attn_factor, corr_dims, theta_scale,
freq_factors, row_indices, set_rows_stride, is_neox);
} else {
ggml_cuda_kernel_launch(rms_norm_mul_rope_f32<256, true, D>, launch_params,
x, dst, ncols, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
x, dst, ncols, nchannels, nsamples, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
n_dims, pos, freq_scale, ext_factor, attn_factor, corr_dims, theta_scale,
freq_factors, row_indices, set_rows_stride, is_neox);
@@ -836,13 +846,13 @@ static void rms_norm_mul_rope_cuda(
const ggml_cuda_kernel_launch_params launch_params = {blocks_num, block_dims, 32*sizeof(float), stream};
if (freq_factors == nullptr) {
ggml_cuda_kernel_launch(rms_norm_mul_rope_f32<1024, false, D>, launch_params,
x, dst, ncols, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
x, dst, ncols, nchannels, nsamples, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
n_dims, pos, freq_scale, ext_factor, attn_factor, corr_dims, theta_scale,
freq_factors, row_indices, set_rows_stride, is_neox);
} else {
ggml_cuda_kernel_launch(rms_norm_mul_rope_f32<1024, true, D>, launch_params,
x, dst, ncols, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
x, dst, ncols, nchannels, nsamples, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
n_dims, pos, freq_scale, ext_factor, attn_factor, corr_dims, theta_scale,
freq_factors, row_indices, set_rows_stride, is_neox);
+124 -66
View File
@@ -1,6 +1,29 @@
#include "argsort.cuh"
#include "top-k.cuh"
// Adjusted implementation thresholds from #28547, can be overridden at build time
#ifndef GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC
# if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA)
// not measured on HIP/MUSA, keep the old split
# define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC 1024
# else
# define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC 512
# endif
#endif // GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC
#ifndef GGML_CUDA_TOP_K_NCOLS_THRESHOLD_ARGSORT
# define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_ARGSORT 4096
#endif // GGML_CUDA_TOP_K_NCOLS_THRESHOLD_ARGSORT
// bitonic up to this width while nrows fits in one wave of SMs, 0 disables
#ifndef GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS
# if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA)
# define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS 0
# else
# define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS 1024
# endif
#endif // GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS
#ifdef GGML_CUDA_USE_CUB
# include <cub/cub.cuh>
// DeviceTopK has a race condition before CCCL 3.4.3.
@@ -14,6 +37,15 @@ using namespace cub;
# endif // CCCL >= 3.4.3
#endif // GGML_CUDA_USE_CUB
// max rows for the per-row DeviceTopK / CUB argsort path before switching to radix / bitonic
#ifndef GGML_CUDA_TOP_K_NROWS_THRESHOLD
# ifdef CUB_TOP_K_AVAILABLE
# define GGML_CUDA_TOP_K_NROWS_THRESHOLD 2
# else
# define GGML_CUDA_TOP_K_NROWS_THRESHOLD 1
# endif
#endif // GGML_CUDA_TOP_K_NROWS_THRESHOLD
#ifdef CUB_TOP_K_AVAILABLE
static void top_k_cub(ggml_cuda_pool & pool,
@@ -40,7 +72,7 @@ static void top_k_cub(ggml_cuda_pool & pool,
ncols, k, env));
}
#elif defined(GGML_CUDA_USE_CUB) // CUB_TOP_K_AVAILABLE
#endif // CUB_TOP_K_AVAILABLE
static int next_power_of_2(int x) {
int n = 1;
@@ -50,10 +82,6 @@ static int next_power_of_2(int x) {
return n;
}
#endif // CUB_TOP_K_AVAILABLE
#if !defined(GGML_CUDA_USE_CUB) && defined(GGML_USE_HIP)
static __device__ __forceinline__ uint32_t top_k_float_to_ordered(float value) {
const uint32_t bits = __float_as_uint(value);
const uint32_t mask = (uint32_t) (-(int32_t) (bits >> 31)) | 0x80000000U;
@@ -95,7 +123,7 @@ static __global__ void top_k_radix_histogram(
__syncthreads();
const top_k_radix_state state = states[row];
for (int col = row_block * BLOCK_SIZE + tid;
for (int64_t col = row_block * BLOCK_SIZE + tid;
col < ncols;
col += blocks_per_row * BLOCK_SIZE) {
const uint32_t key = top_k_float_to_ordered(row_src[col]);
@@ -165,7 +193,7 @@ static __global__ void top_k_radix_gather(
int * row_dst = dst + (size_t) row * k;
top_k_radix_state * state = &states[row];
for (int col = row_block * BLOCK_SIZE + tid;
for (int64_t col = row_block * BLOCK_SIZE + tid;
col < ncols;
col += blocks_per_row * BLOCK_SIZE) {
const uint32_t key = top_k_float_to_ordered(row_src[col]);
@@ -183,36 +211,72 @@ static __global__ void top_k_radix_gather(
static void top_k_radix_cuda(
ggml_cuda_pool & pool,
const float * src, int * dst, int ncols, int nrows, int k, cudaStream_t stream) {
const float * src, int * dst, int ncols, int64_t nrows, int k, cudaStream_t stream) {
constexpr int BLOCK_SIZE = 256;
constexpr int RADIX_BITS = 8;
constexpr int NBINS = 1 << RADIX_BITS;
const int blocks_per_row = std::min((ncols + 1023) / 1024, 64);
const int blocks_per_row = (int) std::min<int64_t>(((int64_t) ncols + 1023) / 1024, 64);
ggml_cuda_pool_alloc<top_k_radix_state> states_alloc(pool, nrows);
ggml_cuda_pool_alloc<int> histograms_alloc(pool, (size_t) nrows * blocks_per_row * NBINS);
// chunk the rows to bound the histogram memory to 64 MB
const int64_t chunk_nrows = ggml_cuda_chunk_nrows((size_t) blocks_per_row * NBINS * sizeof(int), nrows);
ggml_cuda_pool_alloc<top_k_radix_state> states_alloc(pool, chunk_nrows);
ggml_cuda_pool_alloc<int> histograms_alloc(pool, (size_t) chunk_nrows * blocks_per_row * NBINS);
top_k_radix_state * states = states_alloc.get();
int * histograms = histograms_alloc.get();
top_k_radix_init<<<(nrows + BLOCK_SIZE - 1) / BLOCK_SIZE, BLOCK_SIZE, 0, stream>>>(states, nrows, k);
for (int64_t i = 0; i < nrows; i += chunk_nrows) {
const int iter_nrows = std::min(chunk_nrows, nrows - i);
const dim3 row_grid(blocks_per_row * nrows);
for (int shift = 32 - RADIX_BITS; shift >= 0; shift -= RADIX_BITS) {
top_k_radix_histogram<BLOCK_SIZE, RADIX_BITS>
top_k_radix_init<<<(iter_nrows + BLOCK_SIZE - 1) / BLOCK_SIZE, BLOCK_SIZE, 0, stream>>>(states, iter_nrows, k);
const dim3 row_grid(blocks_per_row * iter_nrows);
for (int shift = 32 - RADIX_BITS; shift >= 0; shift -= RADIX_BITS) {
top_k_radix_histogram<BLOCK_SIZE, RADIX_BITS>
<<<row_grid, BLOCK_SIZE, 0, stream>>>(
src, states, histograms, ncols, blocks_per_row, shift);
top_k_radix_select<BLOCK_SIZE, RADIX_BITS>
<<<iter_nrows, BLOCK_SIZE, 0, stream>>>(histograms, states, blocks_per_row, shift);
}
top_k_radix_reset_counters
<<<(iter_nrows + BLOCK_SIZE - 1) / BLOCK_SIZE, BLOCK_SIZE, 0, stream>>>(states, iter_nrows);
top_k_radix_gather<BLOCK_SIZE>
<<<row_grid, BLOCK_SIZE, 0, stream>>>(
src, states, histograms, ncols, blocks_per_row, shift);
top_k_radix_select<BLOCK_SIZE, RADIX_BITS>
<<<nrows, BLOCK_SIZE, 0, stream>>>(histograms, states, blocks_per_row, shift);
}
src, dst, states, ncols, k, blocks_per_row);
top_k_radix_reset_counters
<<<(nrows + BLOCK_SIZE - 1) / BLOCK_SIZE, BLOCK_SIZE, 0, stream>>>(states, nrows);
top_k_radix_gather<BLOCK_SIZE>
<<<row_grid, BLOCK_SIZE, 0, stream>>>(
src, dst, states, ncols, k, blocks_per_row);
src += (size_t) ncols * iter_nrows;
dst += (size_t) k * iter_nrows;
}
}
#endif // !defined(GGML_CUDA_USE_CUB) && defined(GGML_USE_HIP)
static void top_k_argsort_cuda(
ggml_cuda_pool & pool,
const float * src, int * dst, int ncols, int64_t nrows, int k, bool use_cub, cudaStream_t stream) {
const int64_t chunk_nrows = ggml_cuda_chunk_nrows((size_t) ncols * sizeof(int), nrows);
ggml_cuda_pool_alloc<int> tmp_alloc(pool, (size_t) ncols * chunk_nrows);
int * tmp = tmp_alloc.get();
for (int64_t i = 0; i < nrows; i += chunk_nrows) {
const int iter_nrows = std::min(chunk_nrows, nrows - i);
if (use_cub) {
#ifdef GGML_CUDA_USE_CUB
argsort_f32_i32_cuda_cub(pool, src, tmp, ncols, iter_nrows, GGML_SORT_ORDER_DESC, stream);
#else
GGML_ABORT("CUB is not available");
#endif // GGML_CUDA_USE_CUB
} else {
argsort_f32_i32_cuda_bitonic(src, tmp, ncols, iter_nrows, GGML_SORT_ORDER_DESC, stream);
}
CUDA_CHECK(cudaMemcpy2DAsync(dst, k * sizeof(int), tmp, ncols * sizeof(int), k * sizeof(int), iter_nrows,
cudaMemcpyDeviceToDevice, stream));
src += (size_t) ncols * iter_nrows;
dst += (size_t) k * iter_nrows;
}
}
void ggml_cuda_op_top_k(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
const ggml_tensor * src0 = dst->src[0];
@@ -229,51 +293,45 @@ void ggml_cuda_op_top_k(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
const int64_t nrows = ggml_nrows(src0);
const int64_t k = dst->ne[0];
ggml_cuda_pool & pool = ctx.pool();
const int device = ggml_cuda_get_device();
#ifdef CUB_TOP_K_AVAILABLE
// TODO: Switch to `DeviceSegmentedTopK` for multi-row TopK once implemented
// https://github.com/NVIDIA/cccl/issues/6391
// TODO: investigate if there exists a point where parallelized argsort is faster than sequential top-k
for (int i = 0; i < nrows; i++) {
// a single row always uses DeviceTopK if available
const bool bitonic_short = nrows > 1 && ncols <= GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC;
#else
const bool bitonic_short = ncols <= GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC;
#endif // CUB_TOP_K_AVAILABLE
const bool bitonic_few_rows = nrows > GGML_CUDA_TOP_K_NROWS_THRESHOLD &&
ncols <= GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS &&
nrows <= ggml_cuda_info().devices[device].nsm;
if (bitonic_short || bitonic_few_rows) {
// the padded row must fit in shared memory
const int ncols_pad = next_power_of_2(ncols);
if (ncols_pad * sizeof(int) <= ggml_cuda_info().devices[device].smpb) {
top_k_argsort_cuda(pool, src0_d, dst_d, ncols, nrows, k, false, stream);
return;
}
}
if (nrows > GGML_CUDA_TOP_K_NROWS_THRESHOLD) {
top_k_radix_cuda(pool, src0_d, dst_d, ncols, nrows, k, stream);
return;
}
#ifdef CUB_TOP_K_AVAILABLE
// TODO: Assess perf of `DeviceBatchedTopK` for multi-row TopK & CCCL >= 3.5.0, re-running perf sweep of https://github.com/ggml-org/llama.cpp/pull/28713
for (int64_t i = 0; i < nrows; i++) {
top_k_cub(pool, src0_d + i * ncols, dst_d + i * k, ncols, k, stream);
}
#elif defined(GGML_CUDA_USE_CUB) // CUB_TOP_K_AVAILABLE
// Fall back to argsort + copy
const int ncols_pad = next_power_of_2(ncols);
const size_t shared_mem = ncols_pad * sizeof(int);
const size_t max_shared_mem = ggml_cuda_info().devices[ggml_cuda_get_device()].smpb;
const bool use_bitonic = shared_mem <= max_shared_mem && ncols <= 1024;
const int chunk_nrows = argsort_f32_i32_cuda_cub_chunk_nrows(src0->nb[1], nrows);
ggml_cuda_pool_alloc<int> temp_dst_alloc(pool, ncols * chunk_nrows);
int * tmp_dst = temp_dst_alloc.get();
for (int64_t i = 0; i < nrows; i += chunk_nrows) {
int iter_nrows = std::min((int64_t) chunk_nrows, nrows - i);
if (use_bitonic) {
argsort_f32_i32_cuda_bitonic(src0_d, tmp_dst, ncols, iter_nrows, GGML_SORT_ORDER_DESC, stream);
} else {
argsort_f32_i32_cuda_cub(pool, src0_d, tmp_dst, ncols, iter_nrows, GGML_SORT_ORDER_DESC, stream);
}
CUDA_CHECK(cudaMemcpy2DAsync(dst_d, k * sizeof(int), tmp_dst, ncols * sizeof(int), k * sizeof(int), iter_nrows,
cudaMemcpyDeviceToDevice, stream));
src0_d += ncols * iter_nrows;
dst_d += k * iter_nrows;
if (ncols <= GGML_CUDA_TOP_K_NCOLS_THRESHOLD_ARGSORT) {
top_k_argsort_cuda(pool, src0_d, dst_d, ncols, nrows, k, true, stream);
} else {
top_k_radix_cuda(pool, src0_d, dst_d, ncols, nrows, k, stream);
}
#else // GGML_CUDA_USE_CUB
#if defined(GGML_USE_HIP)
if (ncols > 1024) {
top_k_radix_cuda(pool, src0_d, dst_d, ncols, nrows, k, stream);
} else {
#endif // defined(GGML_USE_HIP)
ggml_cuda_pool_alloc<int> temp_dst_alloc(pool, ncols * nrows);
int * tmp_dst = temp_dst_alloc.get();
argsort_f32_i32_cuda_bitonic(src0_d, tmp_dst, ncols, nrows, GGML_SORT_ORDER_DESC, stream);
CUDA_CHECK(cudaMemcpy2DAsync(dst_d, k * sizeof(int), tmp_dst, ncols * sizeof(int), k * sizeof(int), nrows,
cudaMemcpyDeviceToDevice, stream));
#if defined(GGML_USE_HIP)
}
#endif // defined(GGML_USE_HIP)
#endif
top_k_radix_cuda(pool, src0_d, dst_d, ncols, nrows, k, stream);
#endif // CUB_TOP_K_AVAILABLE
}
+50 -7
View File
@@ -134,24 +134,67 @@ static void unary_cuda(const T * x, T * dst, const int k, cudaStream_t stream) {
ggml_cuda_kernel_launch(unary_op_kernel<op, T>, launch_params, x, dst, k);
}
template <float (*op)(float), typename T>
static __global__ void unary_op_kernel_strided(const T * x, T * dst, const int k,const int64_t ne00,const int64_t ne01,const int64_t ne02,const size_t nb00,const size_t nb01,const size_t nb02,const size_t nb03) {
ggml_cuda_pdl_lc();
const int i = blockDim.x*blockIdx.x + threadIdx.x;
if (i >= k) {
return;
}
int64_t rem = i;
const int64_t i0 = rem % ne00; rem /= ne00;
const int64_t i1 = rem % ne01; rem /= ne01;
const int64_t i2 = rem % ne02;
const int64_t i3 = rem / ne02;
const size_t src_byte_offset = i0 * nb00 + i1 * nb01 + i2 * nb02 + i3 * nb03;
const T * src_ptr = (const T *)((const char *)x + src_byte_offset);
ggml_cuda_pdl_sync();
dst[i] = ggml_cuda_cast<T>(op(ggml_cuda_cast<float>(*src_ptr)));
}
template <float (*op)(float), typename T>
static void unary_cuda_strided(const T * x, T * dst, const int k,const int64_t ne00,const int64_t ne01,const int64_t ne02,const size_t nb00,const size_t nb01,const size_t nb02,const size_t nb03, cudaStream_t stream) {
const int num_blocks = (k + CUDA_NEG_BLOCK_SIZE - 1) / CUDA_NEG_BLOCK_SIZE;
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params((dim3)num_blocks, CUDA_NEG_BLOCK_SIZE, 0, stream);
ggml_cuda_kernel_launch(unary_op_kernel_strided<op, T>, launch_params, x, dst, k, ne00,ne01,ne02,nb00,nb01,nb02,nb03);
}
template <float (*op)(float)>
void ggml_cuda_op_unary(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
const ggml_tensor * src0 = dst->src[0];
const void * src0_d = src0->data;
void * dst_d = dst->data;
cudaStream_t stream = ctx.stream();
GGML_ASSERT(ggml_is_contiguous(src0));
cudaStream_t stream = ctx.stream();
GGML_ASSERT(src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16 || src0->type == GGML_TYPE_BF16);
GGML_ASSERT(src0->type == dst->type);
if (src0->type == GGML_TYPE_F16) {
unary_cuda<op>((const half *)src0_d, (half *)dst_d, ggml_nelements(src0), stream);
} else if (src0->type == GGML_TYPE_BF16) {
unary_cuda<op>((const nv_bfloat16 *)src0_d, (nv_bfloat16 *)dst_d, ggml_nelements(src0), stream);
if (ggml_is_contiguous(src0)) {
if (src0->type == GGML_TYPE_F16) {
unary_cuda<op>((const half *)src0_d, (half *)dst_d, ggml_nelements(src0), stream);
} else if (src0->type == GGML_TYPE_BF16) {
unary_cuda<op>((const nv_bfloat16 *)src0_d, (nv_bfloat16 *)dst_d, ggml_nelements(src0), stream);
} else {
unary_cuda<op>((const float *)src0_d, (float *)dst_d, ggml_nelements(src0), stream);
}
} else {
unary_cuda<op>((const float *)src0_d, (float *)dst_d, ggml_nelements(src0), stream);
if (src0->type == GGML_TYPE_F16) {
unary_cuda_strided<op>((const half *)src0_d, (half *)dst_d, ggml_nelements(src0),
src0->ne[0], src0->ne[1], src0->ne[2],
src0->nb[0], src0->nb[1], src0->nb[2], src0->nb[3], stream);
} else if (src0->type == GGML_TYPE_BF16) {
unary_cuda_strided<op>((const nv_bfloat16 *)src0_d, (nv_bfloat16 *)dst_d, ggml_nelements(src0),
src0->ne[0], src0->ne[1], src0->ne[2],
src0->nb[0], src0->nb[1], src0->nb[2], src0->nb[3], stream);
} else {
unary_cuda_strided<op>((const float *)src0_d, (float *)dst_d, ggml_nelements(src0),
src0->ne[0], src0->ne[1], src0->ne[2],
src0->nb[0], src0->nb[1], src0->nb[2], src0->nb[3], stream);
}
}
}
+245 -24
View File
@@ -53,6 +53,7 @@
#define GGML_COMMON_IMPL_CPP
#include "ggml-backend-impl.h"
#include "ggml-alloc.h"
#include "ggml-common.h"
#include "ggml-hexagon.h"
#include "ggml-impl.h"
@@ -101,13 +102,16 @@ static size_t opt_ndev = 1;
static size_t opt_nhvx = 0; // use all
static int opt_nhmx = 1; // when set, enable HMX; when 0, use HVX only
static size_t opt_vmem = HTP_OP_MAX_VMEM_DEFAULT; // max available va space for buffer mappings
static size_t opt_mbuf = 1ul * 1024 * 1024 * 1024; // max buffer size
static int opt_etm = 0;
static int opt_verbose = 0;
static int opt_profile = 0; // profiling mode (0-disabled, 1-basic, 2-pmu)
static bool opt_hostbuf = false;
static bool opt_dma64 = false;
static size_t opt_mbuf_dyn = 512ul * 1024 * 1024; // max dynamic (compute) buffer size
static size_t opt_mbuf_static = 1ul * 1024 * 1024 * 1024; // max static (weight/KV) buffer size
static size_t opt_mbuf_total = 0; // total buffer space limit (0 = unconstrained)
static int opt_mm_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported)
static int opt_fa_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported)
static int opt_fa_head_split = 1; // 1 = partition flash_attn by KV heads in multicore (default on), 0 = token-based (original)
@@ -121,7 +125,7 @@ static int opt_ar_scatter = 1; // 1 = reduce-scatter the fused ALLREDUCE+ADD
static u32vec opt_pmu_evt { 0x3, 0x111, 0x100, 0x105, 0x240, 0x256, 0x7D, 0x8C };
static int opt_opbatch = 1280; // max number of ops in a batch
static int opt_opqueue = 32; // max number of pending batches
static int opt_opqueue = 8; // max number of pending batches
static int opt_optrace = 0; // trace buffer size per thread (0 means default)
static int opt_oppoll = 0; // polling for batch completions
static int opt_opfusion = 1; // enable/disable op fusion
@@ -2820,8 +2824,30 @@ static size_t ggml_backend_hexagon_buffer_type_get_alloc_size(ggml_backend_buffe
GGML_UNUSED(buft);
}
static size_t parse_size(const char * str, size_t default_unit = 1024 * 1024) {
if (!str || str[0] == '\0') {
return 0;
}
char * end = NULL;
double val = strtod(str, &end);
if (val < 0) {
return 0;
}
if (end && *end) {
while (*end == ' ') end++;
if (*end == 'k' || *end == 'K') {
return (size_t) (val * 1024);
} else if (*end == 'm' || *end == 'M') {
return (size_t) (val * 1024 * 1024);
} else if (*end == 'g' || *end == 'G') {
return (size_t) (val * 1024 * 1024 * 1024);
}
}
return (size_t) (val * default_unit);
}
static size_t ggml_backend_hexagon_buffer_type_get_max_size(ggml_backend_buffer_type_t buft) {
return opt_mbuf;
return opt_mbuf_dyn;
GGML_UNUSED(buft);
}
@@ -2835,25 +2861,199 @@ static bool ggml_backend_hexagon_host_buffer_type_is_host(ggml_backend_buffer_ty
GGML_UNUSED(buft);
}
struct ggml_backend_hexagon_alloc_buffer_n_plan_item {
size_t size;
int first;
int last;
};
using ggml_backend_hexagon_alloc_buffer_n_plan_t = std::vector<ggml_backend_hexagon_alloc_buffer_n_plan_item>;
static const char * ggml_hexagon_kv_layer_suffix(const struct ggml_tensor * t) {
if (strncmp(t->name, "cache_", 6) != 0) {
return NULL;
}
const char * p = strstr(t->name, "_l");
if (!p || !isdigit((unsigned char)p[2])) {
return NULL;
}
return p;
}
struct ggml_backend_hexagon_alloc_unit {
size_t size;
int first;
int last;
};
static ggml_backend_hexagon_alloc_buffer_n_plan_t ggml_backend_hexagon_alloc_buffer_n_plan(
ggml_backend_buffer_type_t buft, struct ggml_tensor ** tensors, int n_tensors) {
ggml_backend_hexagon_alloc_buffer_n_plan_t plan;
const size_t alignment = ggml_backend_buft_get_alignment(buft);
const size_t max_size = opt_mbuf_static > 0 ? opt_mbuf_static : SIZE_MAX;
std::vector<ggml_backend_hexagon_alloc_unit> units;
int i = 0;
while (i < n_tensors) {
struct ggml_tensor * t = tensors[i];
size_t unit_size = 0;
int unit_first = i;
int unit_last = i + 1;
if (t->data == NULL && t->view_src == NULL) {
unit_size += GGML_PAD(ggml_backend_buft_get_alloc_size(buft, t), alignment);
}
const char * layer_suffix = ggml_hexagon_kv_layer_suffix(t);
while (unit_last < n_tensors) {
struct ggml_tensor * next = tensors[unit_last];
if (next->view_src != NULL) {
unit_last++;
continue;
}
if (layer_suffix != NULL) {
const char * next_suffix = ggml_hexagon_kv_layer_suffix(next);
if (next_suffix != NULL && strcmp(layer_suffix, next_suffix) == 0) {
if (next->data == NULL) {
unit_size += GGML_PAD(ggml_backend_buft_get_alloc_size(buft, next), alignment);
}
unit_last++;
continue;
}
}
break;
}
units.push_back({ unit_size, unit_first, unit_last });
i = unit_last;
}
size_t cur_buf_size = 0;
int cur_buf_first = 0;
for (const auto & unit : units) {
if (unit.size == 0) {
continue;
}
if (cur_buf_size > 0 && (cur_buf_size + unit.size) > max_size) {
plan.push_back({ cur_buf_size, cur_buf_first, unit.first });
cur_buf_size = 0;
cur_buf_first = unit.first;
}
cur_buf_size += unit.size;
}
if (cur_buf_size > 0) {
plan.push_back({ cur_buf_size, cur_buf_first, n_tensors });
}
return plan;
}
static ggml_backend_buffer_t ggml_backend_hexagon_buffer_type_alloc_buffer_n(
ggml_backend_buffer_type_t buft, struct ggml_tensor ** tensors, int n_tensors) {
const ggml_backend_hexagon_alloc_buffer_n_plan_t plan = ggml_backend_hexagon_alloc_buffer_n_plan(buft, tensors, n_tensors);
std::vector<ggml_backend_buffer_t> buffers;
buffers.reserve(plan.size());
for (const auto & item : plan) {
ggml_backend_buffer_t buffer = ggml_backend_buft_alloc_buffer(buft, item.size);
if (buffer == NULL) {
GGML_LOG_ERROR("%s: failed to allocate %s buffer of size %zu\n", __func__, ggml_backend_buft_name(buft), item.size);
for (ggml_backend_buffer_t b : buffers) {
ggml_backend_buffer_free(b);
}
return NULL;
}
struct ggml_tallocr tallocr = ggml_tallocr_new(buffer);
struct ggml_tensor * t_failed = NULL;
for (int j = item.first; j < item.last; j++) {
struct ggml_tensor * t = tensors[j];
if (t->data == NULL) {
if (t->view_src == NULL) {
if (ggml_tallocr_alloc(&tallocr, t) != GGML_STATUS_SUCCESS) {
t_failed = t;
break;
}
} else if (t->buffer == NULL) {
if (ggml_backend_view_init(t) != GGML_STATUS_SUCCESS) {
t_failed = t;
break;
}
}
} else {
if (t->view_src != NULL && t->buffer == NULL) {
if (ggml_backend_view_init(t) != GGML_STATUS_SUCCESS) {
t_failed = t;
break;
}
}
}
}
if (t_failed != NULL) {
GGML_LOG_ERROR("%s: failed to initialize tensor %s\n", __func__, t_failed->name);
for (ggml_backend_buffer_t b : buffers) {
ggml_backend_buffer_free(b);
}
ggml_backend_buffer_free(buffer);
return NULL;
}
buffers.push_back(buffer);
}
if (buffers.empty()) {
return NULL;
}
if (buffers.size() == 1) {
return buffers[0];
}
return ggml_backend_multi_buffer_alloc_buffer(buffers.data(), buffers.size());
}
static size_t ggml_backend_hexagon_buffer_type_get_alloc_size_n(
ggml_backend_buffer_type_t buft, struct ggml_tensor ** tensors, int n_tensors) {
const ggml_backend_hexagon_alloc_buffer_n_plan_t plan = ggml_backend_hexagon_alloc_buffer_n_plan(buft, tensors, n_tensors);
size_t total = 0;
for (const auto & item : plan) {
total += item.size;
}
return total;
}
static ggml_backend_buffer_type_i ggml_backend_hexagon_buffer_type_interface = {
/* .get_name = */ ggml_backend_hexagon_buffer_type_name,
/* .alloc_buffer = */ ggml_backend_hexagon_buffer_type_alloc_buffer,
/* .alloc_buffer_n = */ NULL,
/* .alloc_buffer_n = */ ggml_backend_hexagon_buffer_type_alloc_buffer_n,
/* .get_alignment = */ ggml_backend_hexagon_buffer_type_get_alignment,
/* .get_max_size = */ ggml_backend_hexagon_buffer_type_get_max_size,
/* .get_alloc_size = */ ggml_backend_hexagon_buffer_type_get_alloc_size,
/* .get_alloc_size_n = */ NULL,
/* .get_alloc_size_n = */ ggml_backend_hexagon_buffer_type_get_alloc_size_n,
/* .is_host = */ ggml_backend_hexagon_buffer_type_is_host,
};
static ggml_backend_buffer_type_i ggml_backend_hexagon_host_buffer_type_interface = {
/* .get_name = */ ggml_backend_hexagon_buffer_type_name,
/* .alloc_buffer = */ ggml_backend_hexagon_host_buffer_type_alloc_buffer,
/* .alloc_buffer_n = */ NULL,
/* .alloc_buffer_n = */ ggml_backend_hexagon_buffer_type_alloc_buffer_n,
/* .get_alignment = */ ggml_backend_hexagon_buffer_type_get_alignment,
/* .get_max_size = */ ggml_backend_hexagon_buffer_type_get_max_size,
/* .get_alloc_size = */ ggml_backend_hexagon_buffer_type_get_alloc_size,
/* .get_alloc_size_n = */ NULL,
/* .get_alloc_size_n = */ ggml_backend_hexagon_buffer_type_get_alloc_size_n,
/* .is_host = */ ggml_backend_hexagon_host_buffer_type_is_host,
};
@@ -7232,7 +7432,7 @@ static bool ggml_hexagon_supported_get_rows(const struct ggml_hexagon_session *
const struct ggml_tensor * src1 = op->src[1]; // indices
const struct ggml_tensor * dst = op;
if (src0->type == GGML_TYPE_Q4_0 && src0->view_src) {
if ((src0->type == GGML_TYPE_Q4_0 || src0->type == GGML_TYPE_Q4_K || src0->type == GGML_TYPE_Q6_K) && src0->view_src) {
return false;
}
@@ -7241,7 +7441,7 @@ static bool ggml_hexagon_supported_get_rows(const struct ggml_hexagon_session *
if (src0_base->buffer && ggml_backend_buffer_is_hexagon(src0_base->buffer) && src0_base->extra) {
const auto * extra = (const ggml_hexagon_tensor_extra *) src0_base->extra;
is_repacked = (extra->flags & GGML_HEXAGON_TENSOR_REPACK) != 0;
if (is_repacked && src0->type != GGML_TYPE_Q4_0 && src0->type != GGML_TYPE_Q8_0) {
if (is_repacked && src0->type != GGML_TYPE_Q4_0 && src0->type != GGML_TYPE_Q4_K && src0->type != GGML_TYPE_Q6_K && src0->type != GGML_TYPE_Q8_0) {
return false;
}
}
@@ -7252,7 +7452,7 @@ static bool ggml_hexagon_supported_get_rows(const struct ggml_hexagon_session *
return false;
}
if (src0->type == GGML_TYPE_Q4_0 && src0->buffer && !is_repacked) {
if ((src0->type == GGML_TYPE_Q4_0 || src0->type == GGML_TYPE_Q4_K || src0->type == GGML_TYPE_Q6_K) && src0->buffer && ggml_backend_buffer_get_size(src0->buffer) != 0 && !is_repacked) {
return false;
}
@@ -7261,7 +7461,11 @@ static bool ggml_hexagon_supported_get_rows(const struct ggml_hexagon_session *
}
if (src0->type != GGML_TYPE_F32 && src0->type != GGML_TYPE_F16 &&
src0->type != GGML_TYPE_Q4_0 && src0->type != GGML_TYPE_Q8_0 && src0->type != GGML_TYPE_I32) {
src0->type != GGML_TYPE_Q4_0 && src0->type != GGML_TYPE_Q4_K && src0->type != GGML_TYPE_Q6_K && src0->type != GGML_TYPE_Q8_0 && src0->type != GGML_TYPE_I32) {
return false;
}
if ((src0->type == GGML_TYPE_Q4_K || src0->type == GGML_TYPE_Q6_K) && (!ggml_is_contiguous(src0) || ggml_is_permuted(src0) || src0->ne[0] % QK_K)) {
return false;
}
@@ -7290,8 +7494,8 @@ static bool ggml_hexagon_supported_get_rows(const struct ggml_hexagon_session *
return false;
}
// Q4_0 has no raw fallback. Mark only accepted tensors for repacking.
if (src0->type == GGML_TYPE_Q4_0 && !src0->buffer) {
// Tiled quantized weights have no raw fallback. Mark only accepted tensors for repacking.
if ((src0->type == GGML_TYPE_Q4_0 || src0->type == GGML_TYPE_Q4_K || src0->type == GGML_TYPE_Q6_K) && !src0->buffer) {
sess->needs_repack.insert(src0);
}
@@ -7677,7 +7881,7 @@ static htp_op_code op_remap_to_htp(const ggml_tensor * t) {
switch (ggml_get_unary_op(t)) {
case GGML_UNARY_OP_SILU: return HTP_OP_UNARY_SILU;
case GGML_UNARY_OP_GELU: return HTP_OP_UNARY_GELU;
case GGML_UNARY_OP_GELU_QUICK: return HTP_OP_UNARY_GELU;
case GGML_UNARY_OP_GELU_QUICK: return HTP_OP_UNARY_GELU_QUICK;
case GGML_UNARY_OP_GELU_ERF: return HTP_OP_UNARY_GELU_ERF;
case GGML_UNARY_OP_SIGMOID: return HTP_OP_UNARY_SIGMOID;
case GGML_UNARY_OP_NEG: return HTP_OP_UNARY_NEG;
@@ -8541,8 +8745,8 @@ static const char * ggml_backend_hexagon_device_get_description(ggml_backend_dev
}
static void ggml_backend_hexagon_device_get_memory(ggml_backend_dev_t dev, size_t * free, size_t * total) {
*free = 0;
*total = *free;
*free = opt_mbuf_total;
*total = opt_mbuf_total;
GGML_UNUSED(dev);
}
@@ -9291,7 +9495,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
size_t MiB = 1024 * 1024;
// Update vmem default
opt_vmem = opt_arch >= 75 ? HTP_OP_MAX_VMEM_DEFAULT : 3000 * MiB;
opt_vmem = opt_arch >= 75 ? HTP_OP_MAX_VMEM_DEFAULT : 3000 * MiB;
opt_dma64 = opt_arch > 79 && (!str_dma64 || atoi(str_dma64) != 0);
auto RE_ICASE = std::regex_constants::icase;
@@ -9309,13 +9513,30 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
opt_nhmx = str_nhmx ? atoi(str_nhmx) : opt_nhmx;
opt_mm_select = str_mm_select ? atoi(str_mm_select) : opt_mm_select;
opt_fa_select = str_fa_select ? atoi(str_fa_select) : opt_fa_select;
opt_fa_head_split = str_fa_head_split ? atoi(str_fa_head_split) : opt_fa_head_split;
opt_gdn_select = str_gdn_select ? atoi(str_gdn_select) : opt_gdn_select;
opt_ar_select = str_ar_select ? atoi(str_ar_select) : opt_ar_select;
opt_ar_scatter = str_ar_scatter ? atoi(str_ar_scatter) : opt_ar_scatter;
opt_mbuf = str_mbuf ? strtoul(str_mbuf, NULL, 0) * MiB : opt_mbuf;
opt_vmem = str_vmem ? strtoul(str_vmem, NULL, 0) * MiB : opt_vmem;
opt_hostbuf = str_hostbuf ? atoi(str_hostbuf) != 0 : opt_hostbuf;
opt_fa_head_split = str_fa_head_split ? atoi(str_fa_head_split) : opt_fa_head_split;
opt_gdn_select = str_gdn_select ? atoi(str_gdn_select) : opt_gdn_select;
opt_ar_select = str_ar_select ? atoi(str_ar_select) : opt_ar_select;
opt_ar_scatter = str_ar_scatter ? atoi(str_ar_scatter) : opt_ar_scatter;
if (str_mbuf) {
const char * p = str_mbuf;
for (int idx = 0; idx < 3 && p && *p; idx++) {
while (*p == ' ') p++;
const char * comma = strchr(p, ',');
size_t len = comma ? (size_t)(comma - p) : strlen(p);
while (len > 0 && p[len - 1] == ' ') len--;
if (len > 0) {
std::string token(p, len);
if (idx == 0) opt_mbuf_dyn = parse_size(token.c_str());
if (idx == 1) opt_mbuf_static = parse_size(token.c_str());
if (idx == 2) opt_mbuf_total = parse_size(token.c_str());
}
if (!comma) break;
p = comma + 1;
}
}
opt_vmem = str_vmem ? parse_size(str_vmem) : opt_vmem;
opt_hostbuf = str_hostbuf ? atoi(str_hostbuf) != 0 : opt_hostbuf;
// Parse device configuration
const char * str_devices = getenv("GGML_HEXAGON_DEVICES");
+53 -11
View File
@@ -217,7 +217,7 @@ GET_ROWS_THREAD_DT_FN(q8_0, Q8_0_BYTES, int32_t, { compute_get_rows_q8_0((float
GET_ROWS_THREAD_DT_FN(q8_0, Q8_0_BYTES, int64_t, { compute_get_rows_q8_0((float *)dst_spad, src_spad, cur_elems); })
static __attribute__((noinline)) void compute_get_rows_tiled(float * dst, const uint8_t * tile, uint32_t row, bool q4) {
static __attribute__((noinline)) void compute_get_rows_tiled(float * dst, const uint8_t * tile, uint32_t row, bool q4, bool q4_k) {
const HVX_VectorPred first2 = Q6_Q_vsetq_R(2);
const HVX_VectorPred first4 = Q6_Q_vsetq_R(4);
HVX_Vector vq = Q6_V_vzero();
@@ -235,7 +235,9 @@ static __attribute__((noinline)) void compute_get_rows_tiled(float * dst, const
const HVX_Vector lo = Q6_V_vand_VV(vq, Q6_Vb_vsplat_R(0x0F));
const HVX_Vector hi = Q6_Vub_vlsr_VubR(vq, 4);
vq = Q6_V_lo_W(Q6_W_vshuff_VVR(hi, lo, -1));
vq = Q6_Vb_vsub_VbVb(vq, Q6_Vb_vsplat_R(8));
if (!q4_k) {
vq = Q6_Vb_vsub_VbVb(vq, Q6_Vb_vsplat_R(8));
}
} else {
for (int group = 7; group >= 0; --group) {
const HVX_Vector v = Q6_V_vror_VR(hvx_vmem(tile + group * VLEN), 2 * row);
@@ -245,14 +247,46 @@ static __attribute__((noinline)) void compute_get_rows_tiled(float * dst, const
}
}
const HVX_Vector scales = hvx_vmem(tile + (q4 ? 512 : 1024));
const HVX_Vector scale_hf = hvx_vec_repl_f16(Q6_V_vror_VR(scales, 2 * row));
const HVX_Vector scale_hf = hvx_vec_repl_f16(Q6_V_vror_VR(scales, (q4_k ? 4 : 2) * row));
const HVX_Vector scale = Q6_V_lo_W(hvx_vec_f16_to_f32(scale_hf));
const HVX_VectorPair p16 = Q6_Wh_vunpack_Vb(vq);
const HVX_VectorPair p32 = Q6_Ww_vunpack_Vh(Q6_V_lo_W(p16));
const HVX_Vector values = hvx_vec_mul_f32_f32(Q6_Vsf_equals_Vw(Q6_V_lo_W(p32)), scale);
HVX_Vector values = hvx_vec_mul_f32_f32(Q6_Vsf_equals_Vw(Q6_V_lo_W(p32)), scale);
if (q4_k) {
const HVX_Vector offset_hf = hvx_vec_repl_f16(Q6_V_vror_VR(scales, 4 * row + 2));
const HVX_Vector offset = Q6_V_lo_W(hvx_vec_f16_to_f32(offset_hf));
values = hvx_vec_add_f32_f32(values, offset);
}
*(HVX_Vector *) dst = values;
}
static __attribute__((noinline)) void compute_get_rows_q6_k(float * dst, const uint8_t * tile, uint32_t row) {
const HVX_VectorPred first4 = Q6_Q_vsetq_R(4);
const HVX_VectorPred first16 = Q6_Q_vsetq_R(16 * sizeof(float));
const HVX_Vector mask_0f = Q6_Vb_vsplat_R(0x0F);
const HVX_Vector mask_03 = Q6_Vb_vsplat_R(0x03);
HVX_Vector vq = Q6_V_vzero();
for (int group = 7; group >= 0; --group) {
const HVX_Vector lo_plane = Q6_V_vror_VR(hvx_vmem(tile + (group >> 1) * VLEN), 4 * row);
const HVX_Vector hi_plane = Q6_V_vror_VR(hvx_vmem(tile + 512 + (group >> 2) * VLEN), 4 * row);
const HVX_Vector lo = (group & 1) ? Q6_Vub_vlsr_VubR(lo_plane, 4) : Q6_V_vand_VV(lo_plane, mask_0f);
const HVX_Vector hi = Q6_Vub_vlsr_VubR(hi_plane, 2 * (group & 3));
const HVX_Vector packed = Q6_V_vor_VV(lo, Q6_Vw_vasl_VwR(Q6_V_vand_VV(hi, mask_03), 4));
vq = Q6_V_vmux_QVV(first4, packed, Q6_V_vror_VR(vq, VLEN - 4));
}
const HVX_Vector scales = hvx_vmem(tile + 768);
const HVX_Vector scale_lo_hf = hvx_vec_repl_f16(Q6_V_vror_VR(scales, 2 * row));
const HVX_Vector scale_hi_hf = hvx_vec_repl_f16(Q6_V_vror_VR(scales, 64 + 2 * row));
const HVX_Vector scale_lo = Q6_V_lo_W(hvx_vec_f16_to_f32(scale_lo_hf));
const HVX_Vector scale_hi = Q6_V_lo_W(hvx_vec_f16_to_f32(scale_hi_hf));
const HVX_Vector scale = Q6_V_vmux_QVV(first16, scale_lo, scale_hi);
const HVX_VectorPair p16 = Q6_Wh_vunpack_Vb(Q6_Vb_vsub_VbVb(vq, Q6_Vb_vsplat_R(32)));
const HVX_VectorPair p32 = Q6_Ww_vunpack_Vh(Q6_V_lo_W(p16));
*(HVX_Vector *) dst = hvx_vec_mul_f32_f32(Q6_Vsf_equals_Vw(Q6_V_lo_W(p32)), scale);
}
struct get_rows_tiled_task {
dma_addr_t tile_src_base;
dma_addr_t dst_data;
@@ -315,7 +349,9 @@ static void get_rows_thread_tiled(unsigned int nth, unsigned int ith, void * dat
const uint32_t tile_size = grctx->tile_size;
const uint32_t tile_stride = grctx->tile_stride;
const uint32_t dst_bytes = ne00 * sizeof(float);
const bool is_q4 = (octx->src[0]->type == HTP_TYPE_Q4_0);
const bool is_q4 = octx->src[0]->type == HTP_TYPE_Q4_0 || octx->src[0]->type == HTP_TYPE_Q4_K;
const bool is_q4_k = octx->src[0]->type == HTP_TYPE_Q4_K;
const bool is_q6_k = octx->src[0]->type == HTP_TYPE_Q6_K;
for (uint32_t step = 0, spad_idx = 0; step < ir1 - ir0 && spad_idx < 2; ++step, ++spad_idx) {
const uint32_t i = ir0 + step;
@@ -343,7 +379,11 @@ static void get_rows_thread_tiled(unsigned int nth, unsigned int ith, void * dat
for (uint32_t k_tile = 0; k_tile < n_k_tiles; ++k_tile) {
const uint8_t * tile = src_spad + k_tile * tile_stride;
float * dst_block = dst_spad + k_tile * HTP_MM_HMX_TILE_N_COLS;
compute_get_rows_tiled(dst_block, tile, task.row, is_q4);
if (is_q6_k) {
compute_get_rows_q6_k(dst_block, tile, task.row);
} else {
compute_get_rows_tiled(dst_block, tile, task.row, is_q4, is_q4_k);
}
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
@@ -369,10 +409,12 @@ int op_get_rows(struct htp_ops_context * octx) {
const struct htp_get_rows_kernel_params * kparams = (const struct htp_get_rows_kernel_params *) octx->kernel_params;
if (octx->src[0]->type != HTP_TYPE_F32 &&
octx->src[0]->type != HTP_TYPE_F16 &&
octx->src[0]->type != HTP_TYPE_Q4_0 &&
octx->src[0]->type != HTP_TYPE_Q8_0 &&
octx->src[0]->type != HTP_TYPE_I32) {
octx->src[0]->type != HTP_TYPE_F16 &&
octx->src[0]->type != HTP_TYPE_Q4_0 &&
octx->src[0]->type != HTP_TYPE_Q4_K &&
octx->src[0]->type != HTP_TYPE_Q6_K &&
octx->src[0]->type != HTP_TYPE_Q8_0 &&
octx->src[0]->type != HTP_TYPE_I32) {
return HTP_STATUS_NO_SUPPORT;
}
@@ -426,7 +468,7 @@ int op_get_rows(struct htp_ops_context * octx) {
grctx.task_start = task_start;
grctx.tasks = tasks;
grctx.tasks_per_thread = octx->ctx->mdev.count == 1 ? kparams->tasks_per_thread : fastdiv(tasks + n_threads - 1, &octx->n_threads_div);
grctx.tile_size = octx->src[0]->type == HTP_TYPE_Q4_0 ? HTP_MM_WEIGHT_TILE_SIZE_Q4_0 : HTP_MM_WEIGHT_TILE_SIZE_Q8_0;
grctx.tile_size = htp_mm_get_weight_tile_size(octx->src[0]->type);
grctx.tile_stride = (grctx.tile_size + 127) & ~127;
grctx.index_i32 = octx->src[1]->type == HTP_TYPE_I32;
+1 -1
View File
@@ -55,7 +55,7 @@ static inline void htp_get_rows_vtcm_layout_build(
}
if (kernel_type == HTP_GET_ROWS_KERNEL_TILED) {
const size_t tile_size = type == HTP_TYPE_Q4_0 ? HTP_MM_WEIGHT_TILE_SIZE_Q4_0 : HTP_MM_WEIGHT_TILE_SIZE_Q8_0;
const size_t tile_size = htp_mm_get_weight_tile_size(type);
const size_t tile_stride = (tile_size + 127) & ~127;
const uint32_t n_k_tiles = ne00 / HTP_MM_HMX_TILE_N_COLS;
const size_t row_tiles_size = n_k_tiles > 0 ? (n_k_tiles * tile_stride) : tile_stride;
@@ -631,17 +631,40 @@ static void dequantize_tiled_weight_to_fp16_task_q6_k(
HVX_Vector v_scale_k16 = Q6_V_lo_W(Q6_W_vshuff_VVR(v_sc_k16, v_sc_k16, -2));
#pragma unroll
for (int g = 0; g < 8; g++) {
for (int g = 0; g < 8; g += 4) {
const HVX_Vector v_scale = (g < 4) ? v_scale_k0 : v_scale_k16;
HVX_Vector v_q = unpack_q6_k_group(vptr, g, mask_0f, mask_03, i32);
HVX_VectorPair vp16 = Q6_Wh_vunpack_Vb(v_q);
HVX_VectorPair vp_k = Q6_W_vdeal_VVR(Q6_V_hi_W(vp16), Q6_V_lo_W(vp16), -4);
HVX_Vector v_q0 = unpack_q6_k_group(vptr, g + 0, mask_0f, mask_03, i32);
HVX_Vector v_q1 = unpack_q6_k_group(vptr, g + 1, mask_0f, mask_03, i32);
HVX_Vector v_q2 = unpack_q6_k_group(vptr, g + 2, mask_0f, mask_03, i32);
HVX_Vector v_q3 = unpack_q6_k_group(vptr, g + 3, mask_0f, mask_03, i32);
hvx_vmem(dst_ptr + (2 * g + 0) * 64) =
Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_k)), v_scale));
hvx_vmem(dst_ptr + (2 * g + 1) * 64) =
Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_k)), v_scale));
HVX_VectorPair vp16_0 = Q6_Wh_vunpack_Vb(v_q0);
HVX_VectorPair vp16_1 = Q6_Wh_vunpack_Vb(v_q1);
HVX_VectorPair vp16_2 = Q6_Wh_vunpack_Vb(v_q2);
HVX_VectorPair vp16_3 = Q6_Wh_vunpack_Vb(v_q3);
HVX_VectorPair vp_k0 = Q6_W_vdeal_VVR(Q6_V_hi_W(vp16_0), Q6_V_lo_W(vp16_0), -4);
HVX_VectorPair vp_k1 = Q6_W_vdeal_VVR(Q6_V_hi_W(vp16_1), Q6_V_lo_W(vp16_1), -4);
HVX_VectorPair vp_k2 = Q6_W_vdeal_VVR(Q6_V_hi_W(vp16_2), Q6_V_lo_W(vp16_2), -4);
HVX_VectorPair vp_k3 = Q6_W_vdeal_VVR(Q6_V_hi_W(vp16_3), Q6_V_lo_W(vp16_3), -4);
HVX_Vector v_out00 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_k0)), v_scale));
HVX_Vector v_out01 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_k0)), v_scale));
HVX_Vector v_out10 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_k1)), v_scale));
HVX_Vector v_out11 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_k1)), v_scale));
HVX_Vector v_out20 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_k2)), v_scale));
HVX_Vector v_out21 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_k2)), v_scale));
HVX_Vector v_out30 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_k3)), v_scale));
HVX_Vector v_out31 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_k3)), v_scale));
hvx_vmem(dst_ptr + (2 * g + 0) * 64) = v_out00;
hvx_vmem(dst_ptr + (2 * g + 1) * 64) = v_out01;
hvx_vmem(dst_ptr + (2 * g + 2) * 64) = v_out10;
hvx_vmem(dst_ptr + (2 * g + 3) * 64) = v_out11;
hvx_vmem(dst_ptr + (2 * g + 4) * 64) = v_out20;
hvx_vmem(dst_ptr + (2 * g + 5) * 64) = v_out21;
hvx_vmem(dst_ptr + (2 * g + 6) * 64) = v_out30;
hvx_vmem(dst_ptr + (2 * g + 7) * 64) = v_out31;
}
}
}
+1
View File
@@ -117,6 +117,7 @@ enum htp_op_code {
HTP_OP_GLU_GEGLU_ERF,
HTP_OP_POOL_2D,
HTP_OP_POOL_1D,
HTP_OP_UNARY_GELU_QUICK,
HTP_OP_INVALID
};
+2 -1
View File
@@ -12,7 +12,7 @@ static __attribute__((noinline)) HVX_Vector hvx_vec_erf_f32(HVX_Vector x) {
HVX_Vector t = hvx_vec_inverse_f32(hvx_vec_add_f32_f32(
hvx_vec_splat_f32(1.0f), hvx_vec_mul_f32_f32(hvx_vec_splat_f32(0.3275911f), ax)));
HVX_Vector poly = hvx_vec_mul_f32_f32(hvx_vec_splat_f32(1.061405429f), t);
HVX_Vector poly = hvx_vec_splat_f32(1.061405429f);
poly = hvx_vec_add_f32_f32(hvx_vec_splat_f32(-1.453152027f), hvx_vec_mul_f32_f32(poly, t));
poly = hvx_vec_add_f32_f32(hvx_vec_splat_f32(1.421413741f), hvx_vec_mul_f32_f32(poly, t));
poly = hvx_vec_add_f32_f32(hvx_vec_splat_f32(-0.284496736f), hvx_vec_mul_f32_f32(poly, t));
@@ -27,6 +27,7 @@ static __attribute__((noinline)) HVX_Vector hvx_vec_erf_f32(HVX_Vector x) {
return result;
}
// GELU_ERF uses the erf definition.
static inline HVX_Vector hvx_vec_gelu_erf_f32(HVX_Vector x) {
const HVX_Vector scale = hvx_vec_splat_f32(0.7071067811865475f);
const HVX_Vector half = hvx_vec_splat_f32(0.5f);
+7 -4
View File
@@ -351,12 +351,15 @@ IM2COL_BLOCKED_DMA_BODY(im2col_blocked_dma_f32_thread, float, hvx_copy_f32_uu,
const dma_addr_t vsrc = ok \
? (src_data + (size_t) ((in * IC + iic) * IH + iih) * IW * sizeof(float)) \
: src_data; \
dma_queue_push(dma_q, dma_make_data(vdst, vsrc), \
IW * sizeof(float), IW * sizeof(float), IW * sizeof(float), ok ? 1 : 0); \
/* IC*KH descriptors per row can exceed the ring capacity: retire the oldest when full */ \
while (!dma_queue_push(dma_q, dma_make_data(vdst, vsrc), \
IW * sizeof(float), IW * sizeof(float), IW * sizeof(float), \
ok ? 1 : 0)) { \
dma_queue_pop(dma_q); \
} \
} \
} \
for (uint32_t i = 0; i < IC * KH; i++) \
dma_queue_pop(dma_q); \
dma_queue_flush(dma_q); \
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, r); \
for (uint32_t iow = 0; iow < OW; iow++) { \
DST_CTYPE * dst_patch = dstb + (uint64_t) iow * patch_stride; \
+1
View File
@@ -851,6 +851,7 @@ static int execute_op(struct htp_ops_context * octx) {
case HTP_OP_UNARY_SIGMOID:
case HTP_OP_UNARY_SILU:
case HTP_OP_UNARY_GELU:
case HTP_OP_UNARY_GELU_QUICK:
case HTP_OP_UNARY_GELU_ERF:
case HTP_OP_UNARY_NEG:
case HTP_OP_UNARY_EXP:
+2 -2
View File
@@ -425,7 +425,7 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
if (copy_cnt > 0) { \
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct_end); \
if (src2) { \
hvx_add_f32_uaa((uint8_t *) &dst_col[src0_start_row], \
hvx_add_f32_uuu((uint8_t *) &dst_col[src0_start_row], \
(const uint8_t *) tmp, \
(const uint8_t *) ((const float *) mmctx->vtcm_bias + src0_start_row), \
copy_cnt); \
@@ -1108,7 +1108,7 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
if (copy_cnt > 0) {
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, src0_end_row);
if (src2) {
hvx_add_f32_uaa((uint8_t *) &dst_col[src0_start_row],
hvx_add_f32_uuu((uint8_t *) &dst_col[src0_start_row],
(const uint8_t *) tmp,
(const uint8_t *) ((const float *) mmctx->vtcm_bias + src0_start_row),
copy_cnt);
+90 -7
View File
@@ -497,7 +497,72 @@ static void silu_f32(const void * restrict src,
}
}
// gelu(x) = x * sigmoid(1.702 * x) (quick/sigmoid approximation, matches CPU GELU_QUICK reference)
// GELU uses the tanh approximation.
static __attribute__((noinline)) HVX_Vector hvx_vec_gelu_f32(HVX_Vector x) {
const HVX_Vector half = hvx_vec_splat_f32(0.5f);
const HVX_Vector one = hvx_vec_splat_f32(1.0f);
HVX_Vector inner = hvx_vec_mul_f32_f32(x, x);
inner = hvx_vec_mul_f32_f32(inner, hvx_vec_splat_f32(0.044715f));
inner = hvx_vec_add_f32_f32(inner, one);
inner = hvx_vec_mul_f32_f32(inner, x);
inner = hvx_vec_mul_f32_f32(inner, hvx_vec_splat_f32(0.7978845608028654f));
return hvx_vec_mul_f32_f32(hvx_vec_mul_f32_f32(half, x),
hvx_vec_add_f32_f32(one, hvx_vec_tanh_f32(inner)));
}
static inline void hvx_gelu_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) {
assert((unsigned long) dst % 128 == 0);
assert((unsigned long) src % 128 == 0);
HVX_Vector * restrict vdst = (HVX_Vector *) dst;
HVX_Vector * restrict vsrc = (HVX_Vector *) src;
const uint32_t nvec = n / VLEN_FP32;
const uint32_t nloe = n % VLEN_FP32;
uint32_t i = 0;
_Pragma("unroll(4)")
for (; i < nvec; i++) {
vdst[i] = hvx_vec_gelu_f32(vsrc[i]);
}
if (nloe) {
hvx_vec_store_a(&vdst[i], nloe * sizeof(float), hvx_vec_gelu_f32(vsrc[i]));
}
}
// GELU_QUICK uses x * sigmoid(1.702 * x).
static __attribute__((noinline)) HVX_Vector hvx_vec_gelu_quick_f32(HVX_Vector x) {
const HVX_Vector one = hvx_vec_splat_f32(1.0f);
const HVX_Vector max_exp = hvx_vec_splat_f32(87.0f);
const HVX_Vector min_exp = hvx_vec_splat_f32(-87.0f);
const HVX_Vector scaled = hvx_vec_mul_f32_f32(x, hvx_vec_splat_f32(1.702f));
const HVX_Vector sigmoid = hvx_vec_fast_sigmoid_f32_guard(scaled, one, max_exp, min_exp);
return hvx_vec_mul_f32_f32(x, sigmoid);
}
static inline void hvx_gelu_quick_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) {
assert((unsigned long) dst % 128 == 0);
assert((unsigned long) src % 128 == 0);
HVX_Vector * restrict vdst = (HVX_Vector *) dst;
HVX_Vector * restrict vsrc = (HVX_Vector *) src;
const uint32_t nvec = n / VLEN_FP32;
const uint32_t nloe = n % VLEN_FP32;
uint32_t i = 0;
_Pragma("unroll(4)")
for (; i < nvec; i++) {
vdst[i] = hvx_vec_gelu_quick_f32(vsrc[i]);
}
if (nloe) {
hvx_vec_store_a(&vdst[i], nloe * sizeof(float), hvx_vec_gelu_quick_f32(vsrc[i]));
}
}
static void gelu_f32(const void * restrict src,
void * restrict dst,
const uint32_t num_rows,
@@ -508,9 +573,21 @@ static void gelu_f32(const void * restrict src,
const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned);
uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned);
hvx_mul_scalar_f32(dst_local, src_local, 1.702f, ne0);
hvx_sigmoid_f32_aa(dst_local, dst_local, ne0);
hvx_mul_f32_aaa(dst_local, src_local, dst_local, ne0);
hvx_gelu_f32_aa(dst_local, src_local, ne0);
}
}
static void gelu_quick_f32(const void * restrict src,
void * restrict dst,
const uint32_t num_rows,
const struct htp_unary_context * uctx) {
htp_unary_op_preamble;
for (uint32_t ir = 0; ir < num_rows; ir++) {
const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned);
uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned);
hvx_gelu_quick_f32_aa(dst_local, src_local, ne0);
}
}
@@ -774,9 +851,12 @@ static void tile_silu_f32(void * restrict dst, const void * restrict src, uint32
static void tile_gelu_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
(void) uctx;
hvx_mul_scalar_f32((uint8_t *) dst, (const uint8_t *) src, 1.702f, tw);
hvx_sigmoid_f32_aa((uint8_t *) dst, (uint8_t *) dst, tw);
hvx_mul_f32_aaa((uint8_t *) dst, (const uint8_t *) src, (uint8_t *) dst, tw);
hvx_gelu_f32_aa((uint8_t *) dst, (const uint8_t *) src, tw);
}
static void tile_gelu_quick_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
(void) uctx;
hvx_gelu_quick_f32_aa((uint8_t *) dst, (const uint8_t *) src, tw);
}
static void tile_gelu_erf_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
@@ -1533,6 +1613,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
case HTP_OP_UNARY_SIGMOID: op_type = "sigmoid-f32"; break;
case HTP_OP_UNARY_SILU: op_type = "silu-f32"; break;
case HTP_OP_UNARY_GELU: op_type = "gelu-f32"; break;
case HTP_OP_UNARY_GELU_QUICK: op_type = "gelu-quick-f32"; break;
case HTP_OP_UNARY_GELU_ERF: op_type = "gelu-erf-f32"; break;
case HTP_OP_UNARY_SOFTPLUS: op_type = "softplus-f32"; break;
case HTP_OP_UNARY_TANH: op_type = "tanh-f32"; break;
@@ -1684,6 +1765,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
case HTP_OP_UNARY_SIGMOID: compute_func = (void *) tile_sigmoid_f32; break;
case HTP_OP_UNARY_SILU: compute_func = (void *) tile_silu_f32; break;
case HTP_OP_UNARY_GELU: compute_func = (void *) tile_gelu_f32; break;
case HTP_OP_UNARY_GELU_QUICK: compute_func = (void *) tile_gelu_quick_f32; break;
case HTP_OP_UNARY_GELU_ERF: compute_func = (void *) tile_gelu_erf_f32; break;
case HTP_OP_UNARY_SOFTPLUS: compute_func = (void *) tile_softplus_f32; break;
case HTP_OP_UNARY_TANH: compute_func = (void *) tile_tanh_f32; break;
@@ -1731,6 +1813,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
case HTP_OP_UNARY_SIGMOID: compute_func = (void *) sigmoid_f32; break;
case HTP_OP_UNARY_SILU: compute_func = (void *) silu_f32; break;
case HTP_OP_UNARY_GELU: compute_func = (void *) gelu_f32; break;
case HTP_OP_UNARY_GELU_QUICK: compute_func = (void *) gelu_quick_f32; break;
case HTP_OP_UNARY_GELU_ERF: compute_func = (void *) gelu_erf_f32; break;
case HTP_OP_UNARY_SOFTPLUS: compute_func = (void *) softplus_f32; break;
case HTP_OP_UNARY_TANH: compute_func = (void *) tanh_f32; break;
+1
View File
@@ -54,6 +54,7 @@ static inline bool htp_op_is_unary(uint32_t opcode) {
case HTP_OP_UNARY_SIGMOID:
case HTP_OP_UNARY_SILU:
case HTP_OP_UNARY_GELU:
case HTP_OP_UNARY_GELU_QUICK:
case HTP_OP_UNARY_GELU_ERF:
case HTP_OP_UNARY_SOFTPLUS:
case HTP_OP_UNARY_TANH:
+7 -4
View File
@@ -1741,10 +1741,13 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te
op->src[0]->ne[0] != 576) {
return false;
}
if (op->src[1]->ne[0] == 72 && op->src[1]->ne[0] != op->src[2]->ne[0]) {
return false;
}
if (op->src[1]->ne[0] < op->src[2]->ne[0]) {
// the kernels exist for K == V and for these K > V pairs only
if (op->src[1]->ne[0] != op->src[2]->ne[0] &&
!(op->src[1]->ne[0] == 96 && op->src[2]->ne[0] == 64) &&
!(op->src[1]->ne[0] == 128 && op->src[2]->ne[0] == 96) &&
!(op->src[1]->ne[0] == 192 && op->src[2]->ne[0] == 128) &&
!(op->src[1]->ne[0] == 320 && op->src[2]->ne[0] == 256) &&
!(op->src[1]->ne[0] == 576 && op->src[2]->ne[0] == 512)) {
return false;
}
if (op->src[1]->type != op->src[2]->type) {
+1
View File
@@ -3180,6 +3180,7 @@ static int ggml_metal_op_flash_attn_ext_n_kv_max_sparse(const ggml_tensor * op)
(dk == 96 && dv == 96) ||
(dk == 96 && dv == 64) ||
(dk == 128 && dv == 128) ||
(dk == 128 && dv == 96) ||
(dk == 192 && dv == 128) ||
(dk == 192 && dv == 192) ||
(dk == 256 && dv == 256) ||
@@ -40,6 +40,9 @@ int fa_vec_baseline_ne(int dk, int dv) {
if (dk == 128 && dv == 128) {
return 1;
}
if (dk == 128 && dv == 96) {
return 4;
}
if (dk == 192 && dv == 192) {
return 2;
}
+2
View File
@@ -44,6 +44,7 @@ template [[host_name("kernel_flash_attn_ext_f16_dk96_dv96" )]] kernel flash_at
template [[host_name("kernel_flash_attn_ext_f16_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, half4x4, 1, dequantize_f16, half4x4, 1, dequantize_f16, 96, 64>;
template [[host_name("kernel_flash_attn_ext_f16_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, half4x4, 1, dequantize_f16, half4x4, 1, dequantize_f16, 112, 112>;
template [[host_name("kernel_flash_attn_ext_f16_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, half4x4, 1, dequantize_f16, half4x4, 1, dequantize_f16, 128, 128>;
template [[host_name("kernel_flash_attn_ext_f16_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, half4x4, 1, dequantize_f16, half4x4, 1, dequantize_f16, 128, 96>;
template [[host_name("kernel_flash_attn_ext_f16_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, half4x4, 1, dequantize_f16, half4x4, 1, dequantize_f16, 192, 192>;
template [[host_name("kernel_flash_attn_ext_f16_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, half4x4, 1, dequantize_f16, half4x4, 1, dequantize_f16, 192, 128>;
template [[host_name("kernel_flash_attn_ext_f16_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, half4x4, 1, dequantize_f16, half4x4, 1, dequantize_f16, 256, 256>;
@@ -62,6 +63,7 @@ template [[host_name("kernel_flash_attn_ext_bf16_dk96_dv96" )]] kernel flash_at
template [[host_name("kernel_flash_attn_ext_bf16_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_BF, bfloat4x4, 1, dequantize_bf16, bfloat4x4, 1, dequantize_bf16, 96, 64>;
template [[host_name("kernel_flash_attn_ext_bf16_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_BF, bfloat4x4, 1, dequantize_bf16, bfloat4x4, 1, dequantize_bf16, 112, 112>;
template [[host_name("kernel_flash_attn_ext_bf16_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_BF, bfloat4x4, 1, dequantize_bf16, bfloat4x4, 1, dequantize_bf16, 128, 128>;
template [[host_name("kernel_flash_attn_ext_bf16_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_BF, bfloat4x4, 1, dequantize_bf16, bfloat4x4, 1, dequantize_bf16, 128, 96>;
template [[host_name("kernel_flash_attn_ext_bf16_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_BF, bfloat4x4, 1, dequantize_bf16, bfloat4x4, 1, dequantize_bf16, 192, 192>;
template [[host_name("kernel_flash_attn_ext_bf16_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_BF, bfloat4x4, 1, dequantize_bf16, bfloat4x4, 1, dequantize_bf16, 192, 128>;
template [[host_name("kernel_flash_attn_ext_bf16_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_BF, bfloat4x4, 1, dequantize_bf16, bfloat4x4, 1, dequantize_bf16, 256, 256>;
+1
View File
@@ -41,6 +41,7 @@ template [[host_name("kernel_flash_attn_ext_f32_dk96_dv96" )]] kernel flash_at
template [[host_name("kernel_flash_attn_ext_f32_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_F32, float4x4, 1, dequantize_f32, float4x4, 1, dequantize_f32, 96, 64>;
template [[host_name("kernel_flash_attn_ext_f32_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_F32, float4x4, 1, dequantize_f32, float4x4, 1, dequantize_f32, 112, 112>;
template [[host_name("kernel_flash_attn_ext_f32_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_F32, float4x4, 1, dequantize_f32, float4x4, 1, dequantize_f32, 128, 128>;
template [[host_name("kernel_flash_attn_ext_f32_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_F32, float4x4, 1, dequantize_f32, float4x4, 1, dequantize_f32, 128, 96>;
template [[host_name("kernel_flash_attn_ext_f32_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_F32, float4x4, 1, dequantize_f32, float4x4, 1, dequantize_f32, 192, 192>;
template [[host_name("kernel_flash_attn_ext_f32_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_F32, float4x4, 1, dequantize_f32, float4x4, 1, dequantize_f32, 192, 128>;
template [[host_name("kernel_flash_attn_ext_f32_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_F32, float4x4, 1, dequantize_f32, float4x4, 1, dequantize_f32, 256, 256>;
@@ -41,6 +41,7 @@ template [[host_name("kernel_flash_attn_ext_q4_0_dk96_dv96" )]] kernel flash_at
template [[host_name("kernel_flash_attn_ext_q4_0_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_0, 2, dequantize_q4_0, block_q4_0, 2, dequantize_q4_0, 96, 64>;
template [[host_name("kernel_flash_attn_ext_q4_0_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_0, 2, dequantize_q4_0, block_q4_0, 2, dequantize_q4_0, 112, 112>;
template [[host_name("kernel_flash_attn_ext_q4_0_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_0, 2, dequantize_q4_0, block_q4_0, 2, dequantize_q4_0, 128, 128>;
template [[host_name("kernel_flash_attn_ext_q4_0_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_0, 2, dequantize_q4_0, block_q4_0, 2, dequantize_q4_0, 128, 96>;
template [[host_name("kernel_flash_attn_ext_q4_0_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_0, 2, dequantize_q4_0, block_q4_0, 2, dequantize_q4_0, 192, 192>;
template [[host_name("kernel_flash_attn_ext_q4_0_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_0, 2, dequantize_q4_0, block_q4_0, 2, dequantize_q4_0, 192, 128>;
template [[host_name("kernel_flash_attn_ext_q4_0_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_0, 2, dequantize_q4_0, block_q4_0, 2, dequantize_q4_0, 256, 256>;
@@ -41,6 +41,7 @@ template [[host_name("kernel_flash_attn_ext_q4_1_dk96_dv96" )]] kernel flash_at
template [[host_name("kernel_flash_attn_ext_q4_1_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_1, 2, dequantize_q4_1, block_q4_1, 2, dequantize_q4_1, 96, 64>;
template [[host_name("kernel_flash_attn_ext_q4_1_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_1, 2, dequantize_q4_1, block_q4_1, 2, dequantize_q4_1, 112, 112>;
template [[host_name("kernel_flash_attn_ext_q4_1_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_1, 2, dequantize_q4_1, block_q4_1, 2, dequantize_q4_1, 128, 128>;
template [[host_name("kernel_flash_attn_ext_q4_1_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_1, 2, dequantize_q4_1, block_q4_1, 2, dequantize_q4_1, 128, 96>;
template [[host_name("kernel_flash_attn_ext_q4_1_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_1, 2, dequantize_q4_1, block_q4_1, 2, dequantize_q4_1, 192, 192>;
template [[host_name("kernel_flash_attn_ext_q4_1_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_1, 2, dequantize_q4_1, block_q4_1, 2, dequantize_q4_1, 192, 128>;
template [[host_name("kernel_flash_attn_ext_q4_1_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_1, 2, dequantize_q4_1, block_q4_1, 2, dequantize_q4_1, 256, 256>;
@@ -41,6 +41,7 @@ template [[host_name("kernel_flash_attn_ext_q5_0_dk96_dv96" )]] kernel flash_at
template [[host_name("kernel_flash_attn_ext_q5_0_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_0, 2, dequantize_q5_0, block_q5_0, 2, dequantize_q5_0, 96, 64>;
template [[host_name("kernel_flash_attn_ext_q5_0_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_0, 2, dequantize_q5_0, block_q5_0, 2, dequantize_q5_0, 112, 112>;
template [[host_name("kernel_flash_attn_ext_q5_0_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_0, 2, dequantize_q5_0, block_q5_0, 2, dequantize_q5_0, 128, 128>;
template [[host_name("kernel_flash_attn_ext_q5_0_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_0, 2, dequantize_q5_0, block_q5_0, 2, dequantize_q5_0, 128, 96>;
template [[host_name("kernel_flash_attn_ext_q5_0_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_0, 2, dequantize_q5_0, block_q5_0, 2, dequantize_q5_0, 192, 192>;
template [[host_name("kernel_flash_attn_ext_q5_0_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_0, 2, dequantize_q5_0, block_q5_0, 2, dequantize_q5_0, 192, 128>;
template [[host_name("kernel_flash_attn_ext_q5_0_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_0, 2, dequantize_q5_0, block_q5_0, 2, dequantize_q5_0, 256, 256>;
@@ -41,6 +41,7 @@ template [[host_name("kernel_flash_attn_ext_q5_1_dk96_dv96" )]] kernel flash_at
template [[host_name("kernel_flash_attn_ext_q5_1_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_1, 2, dequantize_q5_1, block_q5_1, 2, dequantize_q5_1, 96, 64>;
template [[host_name("kernel_flash_attn_ext_q5_1_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_1, 2, dequantize_q5_1, block_q5_1, 2, dequantize_q5_1, 112, 112>;
template [[host_name("kernel_flash_attn_ext_q5_1_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_1, 2, dequantize_q5_1, block_q5_1, 2, dequantize_q5_1, 128, 128>;
template [[host_name("kernel_flash_attn_ext_q5_1_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_1, 2, dequantize_q5_1, block_q5_1, 2, dequantize_q5_1, 128, 96>;
template [[host_name("kernel_flash_attn_ext_q5_1_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_1, 2, dequantize_q5_1, block_q5_1, 2, dequantize_q5_1, 192, 192>;
template [[host_name("kernel_flash_attn_ext_q5_1_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_1, 2, dequantize_q5_1, block_q5_1, 2, dequantize_q5_1, 192, 128>;
template [[host_name("kernel_flash_attn_ext_q5_1_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_1, 2, dequantize_q5_1, block_q5_1, 2, dequantize_q5_1, 256, 256>;
@@ -41,6 +41,7 @@ template [[host_name("kernel_flash_attn_ext_q8_0_dk96_dv96" )]] kernel flash_at
template [[host_name("kernel_flash_attn_ext_q8_0_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q8_0, 2, dequantize_q8_0, block_q8_0, 2, dequantize_q8_0, 96, 64>;
template [[host_name("kernel_flash_attn_ext_q8_0_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q8_0, 2, dequantize_q8_0, block_q8_0, 2, dequantize_q8_0, 112, 112>;
template [[host_name("kernel_flash_attn_ext_q8_0_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q8_0, 2, dequantize_q8_0, block_q8_0, 2, dequantize_q8_0, 128, 128>;
template [[host_name("kernel_flash_attn_ext_q8_0_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q8_0, 2, dequantize_q8_0, block_q8_0, 2, dequantize_q8_0, 128, 96>;
template [[host_name("kernel_flash_attn_ext_q8_0_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q8_0, 2, dequantize_q8_0, block_q8_0, 2, dequantize_q8_0, 192, 192>;
template [[host_name("kernel_flash_attn_ext_q8_0_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q8_0, 2, dequantize_q8_0, block_q8_0, 2, dequantize_q8_0, 192, 128>;
template [[host_name("kernel_flash_attn_ext_q8_0_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q8_0, 2, dequantize_q8_0, block_q8_0, 2, dequantize_q8_0, 256, 256>;
@@ -56,8 +56,12 @@ template [[host_name("kernel_flash_attn_ext_vec_f16_dk128_dv128_q2_ne4")]] kerne
template [[host_name("kernel_flash_attn_ext_vec_f16_dk128_dv128_q4_ne1")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 128, 128, 1, 4>;
template [[host_name("kernel_flash_attn_ext_vec_f16_dk128_dv128_q4_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 128, 128, 2, 4>;
template [[host_name("kernel_flash_attn_ext_vec_f16_dk128_dv128_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 128, 128, 4, 4>;
template [[host_name("kernel_flash_attn_ext_vec_f16_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 128, 96, 4>;
template [[host_name("kernel_flash_attn_ext_vec_f16_dk128_dv96_q2_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 128, 96, 4, 2>;
template [[host_name("kernel_flash_attn_ext_vec_f16_dk128_dv96_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 128, 96, 4, 4>;
#if defined(GGML_METAL_HAS_BF16)
template [[host_name("kernel_flash_attn_ext_vec_bf16_dk128_dv128")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, bfloat4, 1, dequantize_bf16_t4, bfloat4, 1, dequantize_bf16_t4, 128, 128, 1>;
template [[host_name("kernel_flash_attn_ext_vec_bf16_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, bfloat4, 1, dequantize_bf16_t4, bfloat4, 1, dequantize_bf16_t4, 128, 96, 4>;
#endif
template [[host_name("kernel_flash_attn_ext_vec_f16_dk192_dv192")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 192, 192, 2>;
template [[host_name("kernel_flash_attn_ext_vec_f16_dk192_dv192_q1_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 192, 192, 4, 1>;
@@ -25,6 +25,7 @@ template [[host_name("kernel_flash_attn_ext_vec_f32_dk64_dv64")]] kernel flas
template [[host_name("kernel_flash_attn_ext_vec_f32_dk96_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES_F32, float4, 1, dequantize_f32_t4, float4, 1, dequantize_f32_t4, 96, 96, 4>;
template [[host_name("kernel_flash_attn_ext_vec_f32_dk96_dv64")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES_F32, float4, 1, dequantize_f32_t4, float4, 1, dequantize_f32_t4, 96, 64, 4>;
template [[host_name("kernel_flash_attn_ext_vec_f32_dk128_dv128")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES_F32, float4, 1, dequantize_f32_t4, float4, 1, dequantize_f32_t4, 128, 128, 1>;
template [[host_name("kernel_flash_attn_ext_vec_f32_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES_F32, float4, 1, dequantize_f32_t4, float4, 1, dequantize_f32_t4, 128, 96, 4>;
template [[host_name("kernel_flash_attn_ext_vec_f32_dk192_dv192")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES_F32, float4, 1, dequantize_f32_t4, float4, 1, dequantize_f32_t4, 192, 192, 2>;
template [[host_name("kernel_flash_attn_ext_vec_f32_dk192_dv128")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES_F32, float4, 1, dequantize_f32_t4, float4, 1, dequantize_f32_t4, 192, 128, 2>;
template [[host_name("kernel_flash_attn_ext_vec_f32_dk256_dv256")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES_F32, float4, 1, dequantize_f32_t4, float4, 1, dequantize_f32_t4, 256, 256, 1>;
@@ -44,6 +44,9 @@ template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk128_dv128_q2_ne4")]] kern
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk128_dv128_q4_ne1")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 128, 128, 1, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk128_dv128_q4_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 128, 128, 2, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk128_dv128_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 128, 128, 4, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 128, 96, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk128_dv96_q2_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 128, 96, 4, 2>;
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk128_dv96_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 128, 96, 4, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk192_dv192")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 192, 192, 2>;
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk192_dv192_q1_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 192, 192, 4, 1>;
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk192_dv192_q2_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 192, 192, 2, 2>;
@@ -44,6 +44,9 @@ template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk128_dv128_q2_ne4")]] kern
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk128_dv128_q4_ne1")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 128, 128, 1, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk128_dv128_q4_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 128, 128, 2, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk128_dv128_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 128, 128, 4, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 128, 96, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk128_dv96_q2_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 128, 96, 4, 2>;
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk128_dv96_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 128, 96, 4, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk192_dv192")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 192, 192, 2>;
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk192_dv192_q1_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 192, 192, 4, 1>;
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk192_dv192_q2_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 192, 192, 2, 2>;
@@ -44,6 +44,9 @@ template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk128_dv128_q2_ne4")]] kern
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk128_dv128_q4_ne1")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 128, 128, 1, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk128_dv128_q4_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 128, 128, 2, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk128_dv128_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 128, 128, 4, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 128, 96, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk128_dv96_q2_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 128, 96, 4, 2>;
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk128_dv96_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 128, 96, 4, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk192_dv192")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 192, 192, 2>;
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk192_dv192_q1_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 192, 192, 4, 1>;
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk192_dv192_q2_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 192, 192, 2, 2>;
@@ -44,6 +44,9 @@ template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk128_dv128_q2_ne4")]] kern
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk128_dv128_q4_ne1")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 128, 128, 1, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk128_dv128_q4_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 128, 128, 2, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk128_dv128_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 128, 128, 4, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 128, 96, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk128_dv96_q2_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 128, 96, 4, 2>;
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk128_dv96_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 128, 96, 4, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk192_dv192")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 192, 192, 2>;
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk192_dv192_q1_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 192, 192, 4, 1>;
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk192_dv192_q2_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 192, 192, 2, 2>;
@@ -44,6 +44,9 @@ template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk128_dv128_q2_ne4")]] kern
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk128_dv128_q4_ne1")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 128, 128, 1, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk128_dv128_q4_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 128, 128, 2, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk128_dv128_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 128, 128, 4, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 128, 96, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk128_dv96_q2_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 128, 96, 4, 2>;
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk128_dv96_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 128, 96, 4, 4>;
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk192_dv192")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 192, 192, 2>;
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk192_dv192_q1_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 192, 192, 4, 1>;
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk192_dv192_q2_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 192, 192, 2, 2>;
+4 -1
View File
@@ -21,7 +21,10 @@ if (MUSAToolkit_FOUND)
message(STATUS "MUSA Toolkit found")
if (NOT DEFINED MUSA_ARCHITECTURES)
set(MUSA_ARCHITECTURES "21;22;31")
set(MUSA_ARCHITECTURES "22;31")
endif()
if ("21" IN_LIST MUSA_ARCHITECTURES)
message(WARNING "MUSA architecture 21 (MTT S70, MTT S80, MTT S3000) is deprecated and no longer tested")
endif()
message(STATUS "Using MUSA architectures: ${MUSA_ARCHITECTURES}")
+23
View File
@@ -221,4 +221,27 @@ if (GGML_SYCL_DEVICE_ARCH)
"SHELL:-Xsycl-target-backend=spir64_gen \"-device ${GGML_SYCL_DEVICE_ARCH}\""
-fsycl-max-parallel-link-jobs=${GGML_SYCL_MAX_PARALLEL_LINK_JOBS}
)
# The XMX dequant-GEMM tiles need the sub-group size of the target: 8 on Xe-HPG (DG2, ARL-H),
# 16 on Xe-HPC and Xe2 or newer. ocloc fails on the other size, so build only the one that fits.
# 0 (unknown name, mixed list, or no XMX) builds no XMX tile and the path stays off.
set(_ggml_sycl_xmx_sg "")
string(TOLOWER "${GGML_SYCL_DEVICE_ARCH}" _ggml_sycl_archs)
string(REPLACE "," ";" _ggml_sycl_archs "${_ggml_sycl_archs}")
foreach(_arch IN LISTS _ggml_sycl_archs)
if (_arch MATCHES "^(dg2|acm|ats-m|arl-h|xe-hpg|12\\.5[567]\\.|12\\.74\\.)")
set(_sg 8)
elseif (_arch MATCHES "^(pvc|bmg|lnl|ptl|wcl|nvl|cri|xe2|xe3|xe-hpc|12\\.60\\.|20\\.|30\\.)")
set(_sg 16)
else()
set(_sg 0)
endif()
if (_ggml_sycl_xmx_sg STREQUAL "" OR _ggml_sycl_xmx_sg EQUAL _sg)
set(_ggml_sycl_xmx_sg ${_sg})
else()
set(_ggml_sycl_xmx_sg 0)
endif()
endforeach()
message(STATUS "GGML_SYCL_DEVICE_ARCH: XMX dequant-GEMM sub-group size ${_ggml_sycl_xmx_sg} (0 = off)")
target_compile_definitions(ggml-sycl PRIVATE GGML_SYCL_XMX_AOT_SG=${_ggml_sycl_xmx_sg})
endif()
+62
View File
@@ -65,6 +65,61 @@ extern int g_ggml_sycl_enable_fusion;
extern int g_ggml_sycl_enable_esimd;
extern int g_ggml_sycl_mmvq_wide;
extern int g_ggml_sycl_prioritize_dmmv;
// Which quantized weight formats may take the XMX dequant-GEMM paths. A bitmask rather than one
// flag per path, so a format can be enabled or measured on its own and adding a format is one bit.
enum ggml_sycl_xmx_gather_type {
GGML_SYCL_XMX_GATHER_IQ4_NL = 1 << 0,
GGML_SYCL_XMX_GATHER_IQ3_S = 1 << 1,
GGML_SYCL_XMX_GATHER_IQ4_XS = 1 << 2,
GGML_SYCL_XMX_GATHER_IQ3_XXS = 1 << 3,
GGML_SYCL_XMX_GATHER_IQ2_XXS = 1 << 4,
GGML_SYCL_XMX_GATHER_IQ2_XS = 1 << 5,
GGML_SYCL_XMX_GATHER_IQ2_S = 1 << 6,
GGML_SYCL_XMX_GATHER_IQ1_S = 1 << 7,
GGML_SYCL_XMX_GATHER_IQ1_M = 1 << 8,
GGML_SYCL_XMX_GATHER_Q8_0 = 1 << 9,
GGML_SYCL_XMX_GATHER_Q4_K = 1 << 10,
GGML_SYCL_XMX_GATHER_Q5_K = 1 << 11,
GGML_SYCL_XMX_GATHER_Q6_K = 1 << 12,
};
static constexpr int GGML_SYCL_XMX_GATHER_TYPES_DEFAULT = ~0;
extern int g_ggml_sycl_xmx_gather_types;
// Which joint_matrix combinations the XMX dequant-GEMM paths may use, one bit each (see fused-gemm.cpp).
// GGML_SYCL_DYNAMIC_PRECISION picks the operand type, this mask the combinations of that type.
static constexpr int GGML_SYCL_XMX_GATHER_SHAPES_DEFAULT = 0xff;
extern int g_ggml_sycl_xmx_gather_shapes;
// GGML_SYCL_DYNAMIC_PRECISION: operand type of the XMX dequant-GEMM paths. F32 turns them off and
// keeps the library GEMM in f32. A src1 precision request of an op [TAG_GGML_PREC] is always met.
enum ggml_sycl_dynamic_precision {
GGML_SYCL_DYNAMIC_PRECISION_F16,
GGML_SYCL_DYNAMIC_PRECISION_BF16,
GGML_SYCL_DYNAMIC_PRECISION_TF32,
GGML_SYCL_DYNAMIC_PRECISION_F32,
};
#ifdef GGML_SYCL_F16
static constexpr int GGML_SYCL_DYNAMIC_PRECISION_DEFAULT = GGML_SYCL_DYNAMIC_PRECISION_F16;
#else
static constexpr int GGML_SYCL_DYNAMIC_PRECISION_DEFAULT = GGML_SYCL_DYNAMIC_PRECISION_F32;
#endif
extern int g_ggml_sycl_dynamic_precision;
// GGML_SYCL_DYNAMIC_REQUIRED_PRECISION: the XMX type an F32 src1 request may run on instead of f32
// (TF32, or BF16 which also allows tf32). F32 (default): none. F16: src1 requests are ignored.
extern int g_ggml_sycl_dynamic_required_precision;
// [TAG_GGML_PREC] src1 precision request of the MUL_MAT/MUL_MAT_ID op dst
static inline int32_t ggml_sycl_src1_prec(const ggml_tensor * dst) {
return g_ggml_sycl_dynamic_required_precision == GGML_SYCL_DYNAMIC_PRECISION_F16 ? GGML_PREC_UNDEFINED :
dst->op_params[3];
}
// [TAG_GGML_PREC] the library GEMM and dmmv may convert src1 of the MUL_MAT/MUL_MAT_ID op dst to f16
static inline bool ggml_sycl_src1_f16_ok(const ggml_tensor * dst) {
const int32_t src1_prec = ggml_sycl_src1_prec(dst);
return g_ggml_sycl_dynamic_precision != GGML_SYCL_DYNAMIC_PRECISION_F32 &&
(src1_prec == GGML_PREC_UNDEFINED || src1_prec >= GGML_PREC_F16);
}
extern int g_ggml_sycl_enable_flash_attention;
extern int g_ggml_sycl_dev2dev_memcpy;
extern int g_ggml_sycl_fa_onednn;
@@ -333,6 +388,12 @@ struct mmid_row_mapping {
int32_t i2;
};
struct ggml_sycl_gg_tile {
int32_t expert;
int32_t n0;
int32_t n1;
};
namespace sycl_ex = sycl::ext::oneapi::experimental;
struct ggml_backend_sycl_context {
int device;
@@ -410,6 +471,7 @@ struct ggml_backend_sycl_context {
std::unique_ptr<ggml_sycl_pool> host_pools[GGML_SYCL_MAX_DEVICES];
std::vector<mmid_row_mapping> mmid_row_mapping_host;
std::vector<ggml_sycl_gg_tile> mmid_tile_schedule_host;
static std::unique_ptr<ggml_sycl_pool> new_pool_for_device(queue_ptr qptr, int device);
+88 -7
View File
@@ -452,6 +452,46 @@ static void unary_mul_sycl(const T * x, const T * g, T * dst, const int64_t k, c
});
}
// ADD(bias) + UNARY + MUL(scale) with both broadcast over dim 0, the delta-net alpha gate:
// dst[i] = op(a[i] + bias[i % ne0]) * scale[i % ne0]. k == ne0 makes that the flat index.
template<typename F>
static void add_unary_mul_flat_kernel(const float * a, const float * bias, const float * scale, float * dst,
const int64_t k, const sycl::nd_item<1> &item_ct1, F op) {
SYCL_GLOBAL_ID_LOOP(k, item_ct1) {
dst[i] = op(a[i] + bias[i]) * scale[i];
}
}
template<typename F>
static void add_unary_mul_bcast_kernel(const float * a, const float * bias, const float * scale, float * dst,
const int64_t k, const sycl::uint3 ne0_fd, const sycl::nd_item<1> &item_ct1, F op) {
SYCL_GLOBAL_ID_LOOP(k, item_ct1) {
const uint32_t h = fastmodulo((uint32_t) i, ne0_fd);
dst[i] = op(a[i] + bias[h]) * scale[h];
}
}
template<typename F>
static void add_unary_mul_sycl(const float * a, const float * bias, const float * scale, float * dst,
const int64_t k, const int64_t ne0, queue_ptr main_stream, F op) {
const size_t num_blocks = ceil_div((size_t) k, (size_t) SYCL_GLU_BLOCK_SIZE);
const sycl::nd_range<1> range(num_blocks * sycl::range<1>(SYCL_GLU_BLOCK_SIZE), sycl::range<1>(SYCL_GLU_BLOCK_SIZE));
if (k == ne0) {
main_stream->parallel_for(range, [=](sycl::nd_item<1> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
add_unary_mul_flat_kernel(a, bias, scale, dst, k, item_ct1, op);
});
return;
}
// 32-bit fastdiv, exact only below 2^31; ggml_sycl_can_fuse() already declined past that
GGML_ASSERT(k < ((int64_t) 1 << 31));
const sycl::uint3 ne0_fd = init_fastdiv_values((uint32_t) ne0);
main_stream->parallel_for(range, [=](sycl::nd_item<1> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
add_unary_mul_bcast_kernel(a, bias, scale, dst, k, ne0_fd, item_ct1, op);
});
}
namespace ggml_sycl_detail {
static void acc_f32_sycl(const char *x, const char *y, float *dst,
const int64_t n_elements,
@@ -995,6 +1035,19 @@ static inline void ggml_sycl_op_swiglu(ggml_backend_sycl_context & ctx, ggml_ten
});
}
// Hands `launch` the functor for the unary op of a fused unary chain. Anything else
// ggml_sycl_can_fuse() has already declined, so the default is a dispatcher bug.
template<typename F>
static void dispatch_fused_unary_op(ggml_unary_op uop, F && launch) {
switch (uop) {
case GGML_UNARY_OP_SILU: launch([](float v) { return op_silu(v); }); break;
case GGML_UNARY_OP_SIGMOID: launch([](float v) { return op_sigmoid(v); }); break;
case GGML_UNARY_OP_SOFTPLUS: launch([](float v) { return op_softplus(v); }); break;
default:
GGML_ABORT("fused unary chain: unsupported unary op %s", ggml_unary_op_name(uop));
}
}
// dst = op(unary_node->src[0]) * other, written straight to the MUL output, saving the
// standalone unary launch. Preconditions come from ggml_sycl_can_fuse(); re-asserted here.
void ggml_sycl_op_unary_mul_fused(ggml_backend_sycl_context & ctx, ggml_tensor * unary_node, ggml_tensor * mul_node) {
@@ -1032,13 +1085,41 @@ void ggml_sycl_op_unary_mul_fused(ggml_backend_sycl_context & ctx, ggml_tensor *
}
};
switch (ggml_get_unary_op(unary_node)) {
case GGML_UNARY_OP_SILU: dispatch_type([](float v) { return op_silu(v); }); break;
case GGML_UNARY_OP_SIGMOID: dispatch_type([](float v) { return op_sigmoid(v); }); break;
case GGML_UNARY_OP_SOFTPLUS: dispatch_type([](float v) { return op_softplus(v); }); break;
default:
GGML_ABORT("fused unary+mul: unsupported unary op %s", ggml_unary_op_name(ggml_get_unary_op(unary_node)));
}
dispatch_fused_unary_op(ggml_get_unary_op(unary_node), dispatch_type);
}
// dst = op(a + bias) * scale for an ADD + UNARY + MUL chain whose bias and scale broadcast
// over dim 0. Preconditions come from ggml_sycl_can_fuse(); re-asserted here.
void ggml_sycl_op_add_unary_mul_fused(ggml_backend_sycl_context & ctx, ggml_tensor * add_node,
ggml_tensor * unary_node, ggml_tensor * mul_node) {
// the dst-arity convention the other fusions follow; a and bias live on add_node
scope_op_debug_print scope_dbg_print(__func__, mul_node, /*num_src=*/2);
const ggml_tensor * a = add_node->src[0];
const ggml_tensor * bias = add_node->src[1];
const ggml_tensor * scale = (mul_node->src[0] == unary_node) ? mul_node->src[1] : mul_node->src[0];
// scale is picked by elimination; ggml_can_fuse()'s single-use rule rules out MUL(unary, unary)
GGML_ASSERT(scale != unary_node);
GGML_ASSERT(a->type == GGML_TYPE_F32 && bias->type == GGML_TYPE_F32);
GGML_ASSERT(scale->type == GGML_TYPE_F32 && mul_node->type == GGML_TYPE_F32);
GGML_ASSERT(ggml_are_same_shape(a, mul_node));
// a and dst are indexed flat
GGML_ASSERT(ggml_is_contiguous(a) && ggml_is_contiguous(mul_node));
// bias and scale are one contiguous ne0-length row each, broadcast over the outer dims
GGML_ASSERT(bias->ne[0] == a->ne[0] && scale->ne[0] == a->ne[0]);
GGML_ASSERT(ggml_nrows(bias) == 1 && ggml_nrows(scale) == 1);
GGML_ASSERT(ggml_is_contiguous(bias) && ggml_is_contiguous(scale));
queue_ptr main_stream = ctx.stream();
SYCL_CHECK(ggml_sycl_set_device(ctx.device));
const auto dispatch_op = [&](auto op) {
add_unary_mul_sycl((const float *) a->data, (const float *) bias->data, (const float *) scale->data,
(float *) mul_node->data, ggml_nelements(mul_node), mul_node->ne[0], main_stream, op);
};
dispatch_fused_unary_op(ggml_get_unary_op(unary_node), dispatch_op);
}
__dpct_inline__ float ggml_sycl_op_swiglu_oai_single(float x, float g, float alpha = 1.702f, float limit = 7.0f) {
+5
View File
@@ -132,4 +132,9 @@ void ggml_sycl_arange(ggml_backend_sycl_context & ctx, ggml_tensor * dst);
// fused UNARY(silu|sigmoid|softplus) + MUL; see ggml_sycl_can_fuse() for the accepted shapes
void ggml_sycl_op_unary_mul_fused(ggml_backend_sycl_context & ctx, ggml_tensor * unary_node, ggml_tensor * mul_node);
// fused f32 ADD + UNARY(silu|sigmoid|softplus) + MUL with the bias and the scale broadcast
// over dim 0; see ggml_sycl_can_fuse() for the accepted shapes
void ggml_sycl_op_add_unary_mul_fused(ggml_backend_sycl_context & ctx, ggml_tensor * add_node,
ggml_tensor * unary_node, ggml_tensor * mul_node);
#endif // GGML_SYCL_ELEMENTWISE_HPP

Some files were not shown because too many files have changed in this diff Show More