Compare commits

...
66 Commits
Author SHA1 Message Date
Pascal 23b0202a18 server: leave a busy slot untouched when a request pins it (#30295)
A request asking for a busy id_slot still ran the prompt cache update
on that slot before being deferred. When the RAM cache held a better
match, it was loaded into the slot while another request was still
generating there, and that generation continued on the wrong context.

The busy slot is now returned as is and the request waits for it.
2026-10-10 21:25:38 +02:00
69f201a205 model : support MiniCPM-V 4.7 (#29416)
* mtmd : add MiniCPM-V 4.7 support

Signed-off-by: tc-mb <tianchi_cai@icloud.com>

* model : allow mrope time from an extra position slot

Signed-off-by: tc-mb <tianchi_cai@icloud.com>

* Update conversion/minicpm.py

Co-authored-by: Sigbjørn Skjæret <sigbjorn.skjaeret@huggingface.co>

* Slim down comments

Signed-off-by: tc-mb <tianchi_cai@icloud.com>

* fix for "do not hand-wrap comments"

Signed-off-by: tc-mb <tianchi_cai@icloud.com>

* fix ci

Signed-off-by: tc-mb <tianchi_cai@icloud.com>

* rm 3d repo for pr one

Signed-off-by: tc-mb <tianchi_cai@icloud.com>

* gguf: add rope.section_order metadata

* fix comments

* handle grid layout

* allow compat

---------

Signed-off-by: tc-mb <tianchi_cai@icloud.com>
Co-authored-by: Sigbjørn Skjæret <sigbjorn.skjaeret@huggingface.co>
Co-authored-by: Xuan Son Nguyen <son@huggingface.co>
2026-10-10 20:41:59 +02:00
Aaron Teo abee0c8476 cmake(s390x): disable z17 target for unsupported compilers (#30297) 2026-10-10 20:00:56 +02:00
Mikolaj Kucharski aa94f20861 args: add LLAMA_ARG_SLOT_SAVE_PATH env for --slot-save-path (#30272)
Allow configuring --slot-save-path via LLAMA_ARG_SLOT_SAVE_PATH so
it can be set from an EnvironmentFile in a systemd unit file.
2026-10-10 15:34:40 +02:00
Xuan-Son Nguyen 781dbc5ac9 spec: properly handle mtmd input for mtp (#30257)
* spec: properly handle mtmd input for mtp

* nits
2026-10-10 11:22:38 +02:00
Xuan-Son Nguyen 0fd868cbca mtmd: add build_inp_attn_mask (#30259) 2026-10-10 11:22:12 +02:00
David Friehs 1623d8ce47 cuda: always use MMVQ for MUL_MAT_ID on sm_60 (#27828) 2026-10-10 09:15:42 +03:00
uvos 1bb2b9fcbe CUDA/HIP: fix race in flash_attn_ext_f16_process_tile when nbatch_combine != DKQ/2 (#30103)
Suggested-by: Johannes Gäßler <johannesg@5d6.de>
2026-10-10 09:14:51 +03:00
Aaron Teo 2bbca8f202 ggml-cpu: vectorize fp32 to fp16 conversion (#30157)
ggml-cpu(s390x): rename ulong to uint64_t sized types



ggml-cpu(s390x): rm comment

Signed-off-by: Aaron Teo <aaron.teo1@ibm.com>
2026-10-10 09:08:52 +03:00
Amadeus dyw 404f557b5b vulkan: use 4 rows for NVIDIA MUL_MAT_ID MMVQ (#29274)
Keep the existing RDNA3/4 policy unchanged and update only rm_id to use 4 rows for NVIDIA except pre-Turing.
2026-10-10 09:08:19 +03:00
Aaron Teo b797c82c7d ggml: fix s390x all cpu build (#30140)
Signed-off-by: Aaron Teo <aaron.teo1@ibm.com>
2026-10-10 09:07:52 +03:00
Shawn Gu f2918cabbf opencl: add bin kernels kernel_gemm_moe_q4_k_q8_1_dp4a_bin, kernel_gemm_moe_q6_k_q8_1_dp4a_bin (#30187) 2026-10-09 21:02:11 -07:00
Captain-Tripps 1e6f04a75e sycl : accelerate MXFP4 MoE with arithmetic decoding and weight reordering (#29809) 2026-10-09 22:56:23 -04:00
Xuan-Son Nguyen 10a60cf303 vendor: apply deep nested json patch from upstream (#30253) 2026-10-10 00:07:44 +02:00
Dante 79e2e74eb1 CUDA: fix round issue, under MSVC the CPU and GPU agree (#30229) 2026-10-09 19:59:50 +02:00
Georgi Gerganov 8e2d31e0eb graph : reorder get_rows for embeddings (#30160)
* graph : reorder get_rows for embeddings

* cont : fix gemma4 and improve input embedding construction logic

* cont : add TODO for lora

* cont : fix raw embeddings path

* gemma4 : avoid ple cast in embeddings path
2026-10-09 20:43:05 +03:00
Aleksander GrygierandPascal baef3ed9a1 ui: Models Manager Follow-up Improvements (#30228)
* common : read a GGUF's trained context from its metadata

common_get_gguf_n_ctx_train opens only the file's metadata (no_alloc,
like common_get_decision_type) and reads <arch>.context_length, so a
caller can learn the trained context without loading the model. It
accepts both u32 and u64 values and returns 0 when the file is missing,
unreadable, invalid, or reports no context length.

Assisted-by: pi:zai-org/GLM-5.3-Flash

* server : report the trained context in the models listing

update_caps already resolves the model file offline to read its
modalities, so it now reads the trained context from the same GGUF
metadata, and GET /models reports it as context_length when it is
known. A router listing then carries the context without any Hub
request, which lets the UI sort and filter by it offline.

Assisted-by: pi:zai-org/GLM-5.3-Flash

* ui : take the trained context from the models listing

The router now reports context_length per model, so the option mapping
fills contextLength from it and the manager reads it before the Hub
record. The Context column, the context sort and the context filter
then work with the Hugging Face Hub API turned off. A browser suite
guards the sort and the search, the Hub-cache driven context filter and
the re-sort when details arrive after the sort was clicked.

Assisted-by: pi:zai-org/GLM-5.3-Flash

* ui : mark favorite models with a heart

A favorited model shows a rose heart in the selector even before its
row is hovered, and the crossed heart takes its place on hover, so
unfavoriting stays one hover away. The manager table marks its
favorited rows with the same heart after the badges and capabilities.

Assisted-by: pi:zai-org/GLM-5.3-Flash

* fix: UI text nit

* fix: UI nits

* fix: Favorite models grouping in models table

* feat: Remove sorting from Status column in Models Table

* server: read the GGUF metadata once per model

Read the decision type and the trained context in a single GGUF open,
accept only a UINT32 context length like the model loader, and reset
n_ctx_train with the other caps so a failed refresh drops it.

* fix: Post-review fixes

---------

Co-authored-by: Pascal <admin@serveurperso.com>
2026-10-09 19:28:54 +02:00
Martin Emrich f39148a953 llama-bench: respect -fitc if bigger than required benchmark size (#28331)
Assisted-By: opencode,llama.cpp,Qwen3.6-35B-A3B,Qwen3.8-27B
2026-10-09 19:14:30 +02:00
Anav Prasad 50e3e3e480 CUDA: Remove redundant CUDA copies after SSM_SCAN (#29807)
* CUDA: fuse copy of updated state snapshots into recurrent cache with ssm_scan

* CUDA: remove redundant cuda copies with K==1 (non spec-dec) scenario as well
2026-10-09 18:23:28 +02:00
jingzhou 64df9183f5 opencl: fix kernel compilation for a6x GPUs (#30176)
* opencl: skip kernel_cpy_f32_f32_pack on A6X to avoid shader compiler crash

* The A6x compiler backend found in iot device with a623 (E031.50.31.01)
  cannot handle kernels with a large number of arguments. Skip this
  kernel for A6x to avoid compiler crash

* opencl: A6X constant-fold workaround for get_local_size in GEMV kernels

* opencl: add Adreno 623 to A6X GPU detection list
2026-10-09 09:19:20 -07:00
bosh 6184e92c57 model : use exact GELU for ModernBERT encoders (#30108)
* model : use exact GELU for ModernBERT encoders

Assisted-by: Codex

* model : keep tanh GELU aliases on ggml_geglu

Assisted-by: Claude Opus 5.5

* model : map gelu_python to ggml_geglu_erf

Assisted-by: Claude Opus 5.5
2026-10-09 23:26:56 +08:00
Aldehir Rojas 8b54361025 chat : refactor API (#30210) 2026-10-09 10:06:22 -05:00
Pascal a518119d30 llama: keep the backend sampling graph static across ubatches (#30223)
The reserve builds n_outputs_max_per_seq sampling chains per sampler,
while a decode built one per output row, so the graph changed its
topology after the reserve and GGML_SCHED_NO_REALLOC builds aborted on
the next same sized graph. Every sampler now builds
n_outputs_max_per_seq chains, the ones without a row of the ubatch on
the padding row and not selected, and graph_max_nodes counts them.
2026-10-09 17:20:25 +03:00
Jasmine-tim 8ae386707b ggml: fix OOB write in ggml_acc with negative offset (#30135)
ggml_acc_impl narrowed a size_t offset to int32_t without checking that it
fits, so a large offset could truncate to a negative int32_t. The forward
then sign-extended it to a huge size_t and the bounds assertion wrapped,
allowing an OOB write below the dst buffer. Check the offset before the
narrowing, matching the existing check in ggml_set_impl.
2026-10-09 16:57:33 +03:00
Georgi Gerganov e60eff95fd meta : handle host views (#30217)
* meta : handle views of tensors allocated on the host

A view shares the memory of its view_src, so ggml-alloc never allocates a view in
the buffer of the split it lands in - the scheduler copies the source into the
split and the ops that use the view read that copy. The view node itself is a noop
and does not need a split of its own, but the meta backend asserted when one was
left inside a meta split:

  - ggml_backend_meta_get_split_state() dereferenced tensor->buffer->context
  - the graph rebuild mapped every node with ggml_backend_meta_buffer_simple_tensor()

Accept such nodes when they are views of host tensors, which also generalizes the
previous s_copy_main workaround. This fixes the assert hit by KV cache views when
using --split-mode tensor with partial offload.

Assisted-by: pi:llama.cpp/Qwen3.8-Flash-Next

* archs : re-enable sm tensor for K2 Horizon

* cont : add TODO and reference
2026-10-09 16:01:04 +03:00
Masashi Yoshimura ba6439a6b5 webgpu: use 2D workgroup dispatch for all the ops which use 1D dispatch (e.g., rms_norm) (#30219)
* remove 1D workgroups dispatching

* formatting
2026-10-09 16:00:37 +03:00
Ruben Ortlam 5e4878e978 vulkan: fix rms_norm workgroup count overflow (#30145) 2026-10-09 15:00:14 +02:00
ynankani 609290be6b convert : support compressed-tensor mixed-precision NVFP4 checkpoint (#28636)
Signed-off-by: ynankani <ynankani@nvidia.com>
2026-10-09 13:54:58 +02:00
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
443 changed files with 18136 additions and 5299 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
+3
View File
@@ -45,6 +45,9 @@ insert_final_newline = unset
trim_trailing_whitespace = unset
insert_final_newline = unset
[vendor/**.patch]
trim_trailing_whitespace = unset
[tools/ui/**]
indent_style = unset
indent_size = unset
+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
+3 -3
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");
@@ -3643,7 +3643,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
params.slot_save_path += DIRECTORY_SEPARATOR;
}
}
).set_examples({LLAMA_EXAMPLE_SERVER}));
).set_examples({LLAMA_EXAMPLE_SERVER}).set_env("LLAMA_ARG_SLOT_SAVE_PATH"));
add_opt(common_arg(
{"--media-path"}, "PATH",
"directory for loading local media files; files can be accessed via file:// URLs using relative paths (default: disabled)",
+2 -3
View File
@@ -61,8 +61,7 @@ common_chat_params peg_generator::generate_parser(const common_chat_template &
data.prompt += data.generation_prompt;
}
auto parser = autoparser.build_parser(inputs, parser_generation_prompt);
data.parser = parser.save();
data.parser = autoparser.build_parser(inputs, parser_generation_prompt);
// Build grammar if tools are present
bool has_tools =
@@ -78,7 +77,7 @@ common_chat_params peg_generator::generate_parser(const common_chat_template &
if (include_grammar) {
data.grammar_lazy = !has_response_format && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
// Set grammar triggers based on tool section markers (fall back to per-call markers)
+178 -71
View File
@@ -9,6 +9,7 @@
#include "json.h"
#include "log.h"
#include "parsers/parsers.h"
#include "sampling.h"
#include "jinja/value.h"
#include "jinja/runtime.h"
@@ -112,38 +113,6 @@ const char * common_chat_role_to_string(common_chat_role role) {
return "";
}
json common_chat_msg_delimiters::to_json() const {
json result = json::array();
for (const auto & d : delimiters) {
result.push_back({
{ "role", common_chat_role_to_string(d.role) },
{ "delimiter", d.delimiter },
});
}
return result;
}
common_chat_msg_delimiters common_chat_msg_delimiters_parse(const json & delimiters) {
common_chat_msg_delimiters result;
if (!delimiters.is_array()) {
return result;
}
result.delimiters.reserve(delimiters.size());
for (const auto & d : delimiters) {
if (!d.is_object()) {
continue;
}
result.delimiters.push_back({
common_chat_role_from_string(d.value("role", std::string())),
d.value("delimiter", std::string()),
});
}
return result;
}
void common_chat_msg_delimiters::tokenize(const llama_vocab * vocab) {
for (auto & d : delimiters) {
d.tokens = common_tokenize(vocab, d.delimiter, false, true);
@@ -620,8 +589,11 @@ std::vector<common_chat_tool> common_chat_tools_parse_oaicompat(const json & too
}
common_chat_continuation common_chat_continuation_parse(const common_json & value) {
if (value.is_boolean() && value.get<bool>()) {
return COMMON_CHAT_CONTINUATION_AUTO;
if (value.is_null()) {
return COMMON_CHAT_CONTINUATION_NONE;
}
if (value.is_boolean()) {
return value.get<bool>() ? COMMON_CHAT_CONTINUATION_AUTO : COMMON_CHAT_CONTINUATION_NONE;
}
if (value.is_string()) {
auto value_str = value.get<std::string>();
@@ -632,7 +604,7 @@ common_chat_continuation common_chat_continuation_parse(const common_json & valu
return COMMON_CHAT_CONTINUATION_CONTENT;
}
}
return COMMON_CHAT_CONTINUATION_NONE;
throw std::invalid_argument("Invalid continue_final_message: expected a boolean, \"content\" or \"reasoning_content\"");
}
bool common_chat_verify_template(const std::string & tmpl, bool use_jinja) {
@@ -1087,41 +1059,55 @@ static json common_chat_extra_context() {
return ctx;
}
std::optional<common_chat_params> common_chat_try_specialized_template(
const common_chat_template & tmpl,
const std::string & src,
autoparser::generation_params & params) {
static common_chat_params common_chat_params_init_lfm2_tokens(const common_chat_template & tmpl, const autoparser::generation_params & inputs) {
return common_chat_params_init_lfm2(tmpl, inputs, /* tool_list_tokens = */ true);
}
static common_chat_params common_chat_params_init_lfm2_5(const common_chat_template & tmpl, const autoparser::generation_params & inputs) {
return common_chat_params_init_lfm2(tmpl, inputs, /* tool_list_tokens = */ false);
}
// Older gemma4 templates need their tool responses rewritten before rendering
static common_chat_params common_chat_params_init_gemma4_legacy(const common_chat_template & tmpl, const autoparser::generation_params & inputs) {
auto adjusted = inputs;
workaround::convert_tool_responses_gemma4(adjusted.messages);
return common_chat_params_init_gemma4(tmpl, adjusted);
}
// Pick the dedicated handler for a template from its source, or null for the autoparser.
// Order matters: the first match wins, and later checks assume the earlier ones did not match.
static common_chat_params_init_fn common_chat_template_detect_params_init(const std::string & src) {
// Ministral/Mistral Large 3 - uses special reasoning structure fixes, can't use autoparser
// Note: Mistral Small 3.2 uses [CALL_ID] which Ministral doesn't have, so we can distinguish them
if (src.find("[SYSTEM_PROMPT]") != std::string::npos && src.find("[TOOL_CALLS]") != std::string::npos &&
src.find("[ARGS]") != std::string::npos && src.find("[CALL_ID]") == std::string::npos) {
LOG_DBG("Using specialized template: Ministral/Magistral Large 3\n");
return common_chat_params_init_ministral_3(tmpl, params);
return common_chat_params_init_ministral_3;
}
// LLM-jp-4.1 - GPT-OSS dialect (spaces after special tokens, <|end|>-separated parallel calls)
if (src.find("chat_format=llm-jp-harmony-v1") != std::string::npos) {
LOG_DBG("Using specialized template: LLM-jp Harmony v1\n");
return common_chat_params_init_llm_jp_harmony(tmpl, params);
return common_chat_params_init_llm_jp_harmony;
}
// GPT-OSS - has unique channel-based structure that needs dedicated handler
if (src.find("<|channel|>") != std::string::npos) {
LOG_DBG("Using specialized template: GPT-OSS\n");
return common_chat_params_init_gpt_oss(tmpl, params);
return common_chat_params_init_gpt_oss;
}
// Muse Glimmer format using " to=<recipient>" recipients and <|eom|>/<|eot|> message terminators.
if (src.find("<atem:function_calls>") != std::string::npos && src.find("<|eom|>") != std::string::npos) {
LOG_DBG("Using specialized template: Muse Glimmer\n");
return common_chat_params_init_muse_glimmer(tmpl, params);
return common_chat_params_init_muse_glimmer;
}
// Functionary v3.2 - uses recipient-based format with >>>recipient\n{content}
// Detection: template has ">>>all" for content and ">>>" prefix for tool calls
if (src.find(">>>all") != std::string::npos && src.find(">>>${recipient}") != std::string::npos) {
LOG_DBG("Using specialized template: Functionary v3.2\n");
return common_chat_params_init_functionary_v3_2(tmpl, params);
return common_chat_params_init_functionary_v3_2;
}
// Kimi K2 Thinking - uses unique tool call ID format: functions.<name>:<index>
@@ -1129,14 +1115,14 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
if (src.find("<|tool_calls_section_begin|>") != std::string::npos &&
src.find("<|tool_call_begin|>") != std::string::npos) {
LOG_DBG("Using specialized template: Kimi K2 Thinking\n");
return common_chat_params_init_kimi_k2(tmpl, params);
return common_chat_params_init_kimi_k2;
}
// Kimi K3 - the <|open|>/<|close|>/<|end_of_msg|> markers are unique to it
if (src.find("<|open|>") != std::string::npos && src.find("<|close|>") != std::string::npos &&
src.find("<|end_of_msg|>") != std::string::npos) {
LOG_DBG("Using specialized template: Kimi K3\n");
return common_chat_params_init_kimi_k3(tmpl, params);
return common_chat_params_init_kimi_k3;
}
// K2 Horizon - <|ifm|im_start|> turns, <ifm|think*> reasoning picked by reasoning_effort and
@@ -1144,7 +1130,7 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
if (src.find("<|ifm|im_start|>") != std::string::npos &&
src.find("<ifm|tool_calls>") != std::string::npos) {
LOG_DBG("Using specialized template: K2 Horizon\n");
return common_chat_params_init_k2_horizon(tmpl, params);
return common_chat_params_init_k2_horizon;
}
// Ling 3.0 / Bailing V3 - <role>X</role> sections with <arg_key>/<arg_value> tagged
@@ -1152,7 +1138,7 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
if (src.find("<role>ASSISTANT</role>") != std::string::npos &&
src.find("<arg_key>") != std::string::npos) {
LOG_DBG("Using specialized template: Ling 3.0 (Bailing V3)\n");
return common_chat_params_init_ling3(tmpl, params);
return common_chat_params_init_ling3;
}
// Cohere2 MoE / North Code - marker-wrapped format with <|START_TEXT|> content and
@@ -1161,19 +1147,19 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
if (src.find("<|START_TEXT|>") != std::string::npos &&
src.find("<|START_ACTION|>") != std::string::npos) {
LOG_DBG("Using specialized template: Cohere2 MoE\n");
return common_chat_params_init_cohere2moe(tmpl, params);
return common_chat_params_init_cohere2moe;
}
if (is_lfm2_template(src)) {
LOG_DBG("Using specialized template: LFM2\n");
return common_chat_params_init_lfm2(tmpl, params, /* tool_list_tokens = */ true);
return common_chat_params_init_lfm2_tokens;
}
// LFM2.5 format detection: template uses plain "List of tools: [...]" with no special tokens
if (src.find("List of tools: [") != std::string::npos &&
src.find("<|tool_list_start|>") == std::string::npos) {
LOG_DBG("Using specialized template: LFM2.5\n");
return common_chat_params_init_lfm2(tmpl, params, /* tool_list_tokens = */ false);
return common_chat_params_init_lfm2_5;
}
// GigaChatV3 format detection
@@ -1181,7 +1167,7 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
src.find("<|message_sep|>") != std::string::npos &&
src.find("<|function_call|>") == std::string::npos) {
LOG_DBG("Using specialized template: GigaChatV3\n");
return common_chat_params_init_gigachat_v3(tmpl, params);
return common_chat_params_init_gigachat_v3;
}
// MiniMax-M3: the namespace token "]<]minimax[>[" collides with the autoparser's
@@ -1190,7 +1176,7 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
src.find("<tool_call>") != std::string::npos &&
src.find("<invoke name=") != std::string::npos) {
LOG_DBG("Using specialized template: MiniMax-M3\n");
return common_chat_params_init_minimax_m3(tmpl, params);
return common_chat_params_init_minimax_m3;
}
// DeepSeek V3.2/V4 format detection: template defines dsml_token and uses it for tool calls.
@@ -1201,18 +1187,18 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
(src.find("function_calls") != std::string::npos ||
src.find("tool_calls") != std::string::npos)) {
LOG_DBG("Using specialized template: DeepSeek V3.2/V4\n");
return common_chat_params_init_deepseek_v3_2(tmpl, params);
return common_chat_params_init_deepseek_v3_2;
}
// Gemma4 format detection
if (src.find("'<|tool_call>call:'") != std::string::npos) {
LOG_DBG("Using specialized template: Gemma4\n");
if (src.find("{#- OpenAI Chat Completions:") == std::string::npos) {
// apply workarounds if using the older gemma4 templates
LOG_WRN("%s: detected an outdated gemma4 chat template, applying compatibility workarounds. "
"Consider updating to the official template.\n", __func__);
workaround::convert_tool_responses_gemma4(params.messages);
return common_chat_params_init_gemma4_legacy;
}
return common_chat_params_init_gemma4(tmpl, params);
return common_chat_params_init_gemma4;
}
// MiniCPM5 - XML tool calls with <function name="..."><param name="...">...</param></function>
@@ -1220,14 +1206,14 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
src.find("<function name=\"") != std::string::npos &&
src.find("<param name=\"") != std::string::npos) {
LOG_DBG("Using specialized template: MiniCPM5\n");
return common_chat_params_init_minicpm5(tmpl, params);
return common_chat_params_init_minicpm5;
}
// 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);
return common_chat_params_init_translate_gemma;
}
// Qwen3-Coder XML tool calls, also used by Nemotron Nano 3, Qwen3.5 and StepFun-3.5-Flash
@@ -1237,10 +1223,51 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
// Exclude models that don't use \n between tags
src.find("'<tool_call><function=' ~ tool_call.name ~ '>'") == std::string::npos) {
LOG_DBG("Using specialized template: Qwen3-Coder\n");
return common_chat_params_init_qwen3_coder(tmpl, params);
return common_chat_params_init_qwen3_coder;
}
return std::nullopt;
return nullptr;
}
common_chat_template::common_chat_template(const std::string & src, const std::string & bos_token, const std::string & eos_token) {
jinja::lexer lexer;
auto lexer_res = lexer.tokenize(src);
this->prog = jinja::parse_from_tokens(lexer_res);
this->src = lexer_res.source;
this->bos_tok = bos_token;
this->eos_tok = eos_token;
this->caps = jinja::caps_get(prog);
// LOG_INF("%s: caps:\n%s\n", __func__, this->caps.to_string().c_str());
this->params_init = common_chat_template_detect_params_init(this->src);
if (this->params_init) {
return;
}
// The analysis depends only on the template, so run it once here instead of on every apply.
// A failure is kept for apply to report, so a bad template still loads like it did before.
try {
analysis = std::make_unique<autoparser::autoparser>();
analysis->analyze_template(*this);
} catch (const std::exception & e) {
analysis.reset();
analysis_error = e.what();
}
}
common_chat_template::~common_chat_template() = default;
common_chat_template::common_chat_template(common_chat_template &&) = default;
common_chat_template & common_chat_template::operator=(common_chat_template &&) = default;
std::optional<common_chat_params> common_chat_try_specialized_template(
const common_chat_template & tmpl,
const autoparser::generation_params & params) {
if (!tmpl.params_init) {
return std::nullopt;
}
return tmpl.params_init(tmpl, params);
}
static common_chat_params common_chat_templates_apply_jinja(const struct common_chat_templates * tmpls,
@@ -1342,21 +1369,23 @@ static common_chat_params common_chat_templates_apply_jinja(const struct common_
data.prompt = common_chat_template_direct_apply_impl(tmpl, params_copy);
data.generation_prompt = common_chat_template_generation_prompt_impl(tmpl, params);
data.format = COMMON_CHAT_FORMAT_PEG_NATIVE;
auto parser = build_chat_peg_parser([&data](common_chat_peg_builder &p) {
data.parser = build_chat_peg_parser([&data](common_chat_peg_builder &p) {
return p.literal(data.generation_prompt) << p.content(p.rest());
});
data.parser = parser.save();
return data;
}
if (auto result = common_chat_try_specialized_template(tmpl, src, params)) {
if (auto result = common_chat_try_specialized_template(tmpl, params)) {
return *result;
}
if (!tmpl.analysis) {
throw std::invalid_argument("Unable to generate parser for this template. Automatic parser generation failed: " + tmpl.analysis_error);
}
try {
LOG_DBG("%s: using differential autoparser\n", __func__);
struct autoparser::autoparser autoparser;
autoparser.analyze_template(tmpl);
const auto & autoparser = *tmpl.analysis;
auto auto_params = autoparser::peg_generator::generate_parser(tmpl, params, autoparser);
common_chat_msg_delimiters delimiters;
@@ -1377,8 +1406,7 @@ static common_chat_params common_chat_templates_apply_jinja(const struct common_
auto_params.thinking_end_tags = {std::move(end_tag)};
}
}
common_peg_arena arena;
arena.load(auto_params.parser);
const auto & arena = auto_params.parser;
LOG_DBG("%s: generated parser:\n%s\n\nparser generation prompt: %s\n", __func__, arena.dump(arena.root()).c_str(), auto_params.generation_prompt.c_str());
return auto_params;
} catch (const std::exception & e) {
@@ -1525,9 +1553,10 @@ common_chat_msg common_chat_peg_parse(const common_peg_arena & src_pars
const common_chat_input & input,
bool is_partial,
const common_chat_parser_params & params) {
const common_peg_arena & parser = src_parser.empty() ?
build_chat_peg_parser([](common_chat_peg_builder & p) { return p.content(p.rest()) + p.end(); }) :
src_parser;
// both branches must be lvalues, a temporary here would copy the arena on every call
static const common_peg_arena content_only =
build_chat_peg_parser([](common_chat_peg_builder & p) { return p.content(p.rest()) + p.end(); });
const common_peg_arena & parser = src_parser.empty() ? content_only : src_parser;
if (src_parser.empty()) {
LOG_DBG("No parser definition detected, assuming pure content parser.");
@@ -1598,6 +1627,84 @@ common_chat_msg common_chat_peg_parse(const common_peg_arena & src_pars
return msg;
}
common_chat_session::common_chat_session(const common_chat_templates * tmpls,
const llama_vocab * vocab,
const common_chat_templates_inputs & inputs,
const common_chat_session_params & params) {
auto applied = common_chat_templates_apply(tmpls, inputs);
templated = true;
prompt_text = std::move(applied.prompt);
result.role = "assistant";
grammar_text = std::move(applied.grammar);
grammar_lazy = applied.grammar_lazy;
stops = std::move(applied.additional_stops);
generation_prompt_text = applied.generation_prompt;
thinking_start = std::move(applied.thinking_start_tag);
thinking_ends = std::move(applied.thinking_end_tags);
parser_params.format = applied.format;
parser_params.generation_prompt = vocab ? common_chat_input_tokenize(vocab, applied.generation_prompt)
: common_chat_input(applied.generation_prompt);
parser_params.debug = params.debug;
parser_params.parser = std::move(applied.parser);
delimiters = std::move(applied.message_delimiters);
if (vocab) {
common_params_sampling resolved;
resolved.grammar_lazy = applied.grammar_lazy;
common_sampling_add_preserved_tokens(resolved, vocab, applied.preserved_tokens);
common_sampling_add_grammar_triggers(resolved, vocab, std::move(applied.grammar_triggers));
preserved_tokens = std::move(resolved.preserved_tokens);
grammar_triggers = std::move(resolved.grammar_triggers);
delimiters.tokenize(vocab);
} else {
grammar_triggers = std::move(applied.grammar_triggers);
}
if (inputs.continue_final_message != COMMON_CHAT_CONTINUATION_NONE && !params.echo) {
// start from the prefill so it is not emitted as part of the first delta
result = common_chat_parse(input, true, parser_params);
}
}
void common_chat_session::apply_sampling(common_params_sampling & sampling) const {
if (!templated) {
return;
}
if (!grammar_text.empty()) {
sampling.grammar = {COMMON_GRAMMAR_TYPE_TOOL_CALLS, grammar_text};
}
sampling.grammar_lazy = grammar_lazy;
sampling.generation_prompt = generation_prompt_text;
sampling.preserved_tokens.insert(preserved_tokens.begin(), preserved_tokens.end());
sampling.grammar_triggers.insert(sampling.grammar_triggers.end(), grammar_triggers.begin(), grammar_triggers.end());
}
const common_chat_msg & common_chat_session::feed(const common_chat_input & chunk) {
GGML_ASSERT(!finished && "feed() after finish()");
input.append(chunk);
auto msg = common_chat_parse(input, true, parser_params);
if (!msg.empty()) {
result = std::move(msg);
}
return result;
}
const common_chat_msg & common_chat_session::finish(const common_chat_input & chunk) {
GGML_ASSERT(!finished && "finish() called twice");
finished = true;
input.append(chunk);
auto msg = common_chat_parse(input, false, parser_params);
if (!msg.empty()) {
result = std::move(msg);
}
return result;
}
std::map<std::string, bool> common_chat_templates_get_caps(const common_chat_templates * chat_templates) {
GGML_ASSERT(chat_templates != nullptr);
GGML_ASSERT(chat_templates->template_default != nullptr);
+79 -27
View File
@@ -22,8 +22,16 @@ struct common_chat_templates;
namespace autoparser {
struct generation_params;
struct autoparser;
} // namespace autoparser
struct common_chat_params;
struct common_chat_template;
// Builds the prompt and parser for a template that has a dedicated handler (see common/parsers)
using common_chat_params_init_fn = common_chat_params (*)(const common_chat_template & tmpl,
const autoparser::generation_params & inputs);
struct common_chat_tool_call {
std::string name;
std::string arguments;
@@ -54,19 +62,20 @@ struct common_chat_template {
std::string eos_tok;
std::string src;
chat_template_caps caps;
// Dedicated handler picked once from the source, null when the differential autoparser is used
common_chat_params_init_fn params_init = nullptr;
common_chat_template(const std::string & src, const std::string & bos_token, const std::string & eos_token) {
jinja::lexer lexer;
auto lexer_res = lexer.tokenize(src);
this->prog = jinja::parse_from_tokens(lexer_res);
// Differential analysis, run once here when there is no dedicated handler. Null when there
// is one, or when the analysis failed, in which case analysis_error says why.
std::unique_ptr<autoparser::autoparser> analysis;
std::string analysis_error;
this->src = lexer_res.source;
this->bos_tok = bos_token;
this->eos_tok = eos_token;
common_chat_template(const std::string & src, const std::string & bos_token, const std::string & eos_token);
this->caps = jinja::caps_get(prog);
// LOG_INF("%s: caps:\n%s\n", __func__, this->caps.to_string().c_str());
}
// autoparser is incomplete here, so these are defined where it is complete
~common_chat_template();
common_chat_template(common_chat_template &&);
common_chat_template & operator=(common_chat_template &&);
const std::string & source() const { return src; }
const std::string & bos_token() const { return bos_tok; }
@@ -209,8 +218,6 @@ struct common_chat_msg_delimiters {
// split tokens into message spans. skips maps a start index to a length of a region to jump over without matching
common_chat_msg_spans split(const llama_tokens & tokens, const std::map<size_t, size_t> & skips = {}) const;
common_json to_json() const;
};
struct common_chat_tool {
@@ -278,7 +285,7 @@ struct common_chat_params {
std::vector<common_grammar_trigger> grammar_triggers;
std::vector<std::string> preserved_tokens;
std::vector<std::string> additional_stops;
std::string parser;
common_peg_arena parser;
common_chat_msg_delimiters message_delimiters;
};
@@ -310,16 +317,10 @@ common_chat_input common_chat_input_tokenize(const llama_vocab * vocab, const st
// per-message parsing syntax
// should be derived from common_chat_params
struct common_chat_parser_params {
common_chat_format format = COMMON_CHAT_FORMAT_CONTENT_ONLY;
common_reasoning_format reasoning_format = COMMON_REASONING_FORMAT_NONE; // TODO: refactor this to "bool parse_reasoning"
// Whether reasoning_content should be inlined in the content (e.g. for reasoning_format=deepseek in stream mode)
bool reasoning_in_content = false;
common_chat_input generation_prompt;
bool parse_tool_calls = true;
bool is_continuation = false;
bool echo = false; // Include assistant prefilled msg in output
bool debug = false; // Enable debug output for PEG parser
common_peg_arena parser = {};
common_chat_format format = COMMON_CHAT_FORMAT_CONTENT_ONLY;
common_chat_input generation_prompt;
bool debug = false; // Enable debug output for PEG parser
common_peg_arena parser = {};
common_chat_parser_params() = default;
common_chat_parser_params(const common_chat_params & chat_params) {
format = chat_params.format;
@@ -365,6 +366,60 @@ const char * common_chat_format_name(common_chat_format format);
common_chat_msg common_chat_parse(const common_chat_input & input, bool is_partial, const common_chat_parser_params & params);
common_chat_msg common_chat_peg_parse(const common_peg_arena & src_parser, const common_chat_input & input, bool is_partial, const common_chat_parser_params & params);
struct common_chat_session_params {
bool echo = false; // include the assistant prefill in the output when continuing a message
bool debug = false; // enable debug output for the PEG parser
};
class common_chat_session {
public:
common_chat_session() { result.role = "assistant"; }
common_chat_session(const common_chat_templates * tmpls,
const llama_vocab * vocab,
const common_chat_templates_inputs & inputs,
const common_chat_session_params & params = {});
const std::string & prompt() const { return prompt_text; }
common_chat_format format() const { return parser_params.format; }
const common_chat_msg & msg() const { return result; }
const common_peg_arena & parser() const { return parser_params.parser; }
const std::string & grammar() const { return grammar_text; }
const std::string & generation_prompt() const { return generation_prompt_text; }
const std::string & thinking_start_tag() const { return thinking_start; }
const std::vector<std::string> & thinking_end_tags() const { return thinking_ends; }
const std::vector<std::string> & additional_stops() const { return stops; }
const common_chat_msg_delimiters & message_delimiters() const { return delimiters; }
void apply_sampling(common_params_sampling & sampling) const;
bool has_template() const { return templated; }
const common_chat_msg & feed(const common_chat_input & chunk);
const common_chat_msg & finish(const common_chat_input & chunk = {});
private:
std::string prompt_text;
std::string grammar_text;
bool grammar_lazy = false;
std::vector<common_grammar_trigger> grammar_triggers;
std::set<llama_token> preserved_tokens;
std::vector<std::string> stops;
std::string generation_prompt_text;
std::string thinking_start;
std::vector<std::string> thinking_ends;
common_chat_parser_params parser_params;
common_chat_msg_delimiters delimiters;
common_chat_input input;
common_chat_msg result;
bool templated = false;
bool finished = false;
};
// used by arg and server
const char * common_reasoning_format_name(common_reasoning_format format);
common_reasoning_format common_reasoning_format_from_name(const std::string & format);
@@ -401,8 +456,7 @@ std::string common_chat_template_generation_prompt(
std::optional<common_chat_params> common_chat_try_specialized_template(
const common_chat_template & tmpl,
const std::string & src,
autoparser::generation_params & params);
const autoparser::generation_params & params);
// specialized per-task preset
@@ -412,5 +466,3 @@ struct common_chat_prompt_preset {
};
common_chat_prompt_preset common_chat_get_asr_prompt(const common_chat_templates * chat_templates);
common_chat_msg_delimiters common_chat_msg_delimiters_parse(const common_json & delimiters);
+55 -31
View File
@@ -1171,6 +1171,7 @@ static const std::map<common_decision_type, std::string> COMMON_DECISION_TYPE_NA
{ 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) {
@@ -1194,7 +1195,9 @@ common_decision_type common_get_decision_type(const struct llama_model * model)
return common_decision_type_from_string(buf);
}
common_decision_type common_get_decision_type(const std::string & fname) {
common_gguf_info common_get_gguf_info(const std::string & fname) {
common_gguf_info info;
struct gguf_init_params gguf_params = {
/* .no_alloc = */ true,
/* .ctx = */ nullptr,
@@ -1202,31 +1205,32 @@ common_decision_type common_get_decision_type(const std::string & fname) {
gguf_context_ptr gguf_ctx(gguf_init_from_file(fname.c_str(), gguf_params));
if (!gguf_ctx) {
return COMMON_DECISION_TYPE_UNKNOWN; // missing or unreadable file
return info; // missing or unreadable file
}
std::string arch;
const int64_t arch_id = gguf_find_key(gguf_ctx.get(), "general.architecture");
if (arch_id < 0) {
return COMMON_DECISION_TYPE_UNKNOWN; // no architecture in the metadata
if (arch_id < 0 || gguf_get_kv_type(gguf_ctx.get(), arch_id) != GGUF_TYPE_STRING) {
return info; // no architecture in the metadata
}
if (gguf_get_kv_type(gguf_ctx.get(), arch_id) != GGUF_TYPE_STRING) {
return COMMON_DECISION_TYPE_UNKNOWN; // malformed metadata
}
arch = gguf_get_val_str(gguf_ctx.get(), arch_id);
const std::string arch = gguf_get_val_str(gguf_ctx.get(), arch_id);
if (arch.empty()) {
return COMMON_DECISION_TYPE_UNKNOWN;
return info;
}
const std::string key = arch + ".decision.type";
const int64_t type_id = gguf_find_key(gguf_ctx.get(), key.c_str());
const int64_t type_id = gguf_find_key(gguf_ctx.get(), (arch + ".decision.type").c_str());
if (type_id < 0) {
return COMMON_DECISION_TYPE_NONE;
info.decision_type = COMMON_DECISION_TYPE_NONE;
} else if (gguf_get_kv_type(gguf_ctx.get(), type_id) == GGUF_TYPE_STRING) {
info.decision_type = common_decision_type_from_string(gguf_get_val_str(gguf_ctx.get(), type_id));
}
if (gguf_get_kv_type(gguf_ctx.get(), type_id) != GGUF_TYPE_STRING) {
return COMMON_DECISION_TYPE_UNKNOWN; // malformed metadata
// same key and type as the model loader
const int64_t ctx_id = gguf_find_key(gguf_ctx.get(), (arch + ".context_length").c_str());
if (ctx_id >= 0 && gguf_get_kv_type(gguf_ctx.get(), ctx_id) == GGUF_TYPE_UINT32) {
info.n_ctx_train = gguf_get_val_u32(gguf_ctx.get(), ctx_id);
}
return common_decision_type_from_string(gguf_get_val_str(gguf_ctx.get(), type_id));
return info;
}
common_init_result::common_init_result(common_params & params, bool model_only) :
@@ -1284,7 +1288,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;
@@ -1398,6 +1403,14 @@ std::vector<llama_adapter_lora_ptr> & common_init_result::lora() {
return pimpl->lora;
}
// only for warmup and probe decodes, fill zeros as dummy input
static void common_batch_set_zero_state(common_batch & batch, const llama_model * model, std::vector<float> & zeros) {
zeros.assign(llama_model_n_embd_out(model), 0.0f);
for (int32_t i = 0; i < batch.size(); ++i) {
batch.set_embd_state(i, { zeros.data(), 1, zeros.size() });
}
}
common_init_result_ptr common_init_from_params(common_params & params, bool model_only) {
common_init_result_ptr res(new common_init_result(params, model_only));
@@ -1504,6 +1517,8 @@ common_init_result_ptr common_init_from_params(common_params & params, bool mode
if (llama_model_has_decoder(model)) {
tmp.resize(std::min(tmp.size(), (size_t) params.n_batch));
common_batch batch = common_batch_get_one(lctx, tmp);
std::vector<float> zeros;
common_batch_set_zero_state(batch, model, zeros);
llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
}
llama_memory_clear(llama_get_memory(lctx), true);
@@ -1571,6 +1586,8 @@ common_context_seq_rm_type common_context_can_seq_rm(llama_context * ctx) {
int ret;
{
common_batch batch = common_batch_get_one(ctx, tmp);
std::vector<float> zeros;
common_batch_set_zero_state(batch, llama_get_model(ctx), zeros);
ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
}
if (ret != 0) {
@@ -2156,7 +2173,7 @@ void common_batch::clear() {
}
int32_t common_batch::add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output) {
tokens.push_back({ id, { pos, 0, 0, 0 }, seq_id, output, { nullptr, 0, 0 }, {} });
tokens.push_back({ id, { pos, 0, 0, 0 }, seq_id, output, { nullptr, 0, 0 }, { nullptr, 0, 0 }, {} });
return size() - 1;
}
@@ -2194,8 +2211,16 @@ bool common_batch::set_embd(int32_t idx, llama_embd embd) {
return true;
}
bool common_batch::set_embd_state(int32_t idx, llama_embd state) {
if (idx < 0 || idx >= size() || tokens[idx].state.data != nullptr) {
return false;
}
tokens[idx].state = state;
return true;
}
int32_t common_batch::add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output) {
token t = { LLAMA_TOKEN_NULL, { 0, 0, 0, 0 }, seq_id, output, embd, {} };
token t = { LLAMA_TOKEN_NULL, { 0, 0, 0, 0 }, seq_id, output, embd, { nullptr, 0, 0 }, {} };
for (int32_t j = 0; j < n_pos; ++j) {
t.pos[j] = pos[j];
}
@@ -2240,6 +2265,9 @@ llama_batch_ext * common_batch::get_sub_batch(int32_t off, int32_t n) {
if (t.output) {
llama_batch_ext_set_output_logits(res, idx, true);
}
if (t.state.data) {
llama_batch_ext_set_embd_state(res, idx, t.state); // contexts without a state input ignore it
}
if (t.decision_order != 0) {
llama_batch_ext_set_decision_order(res, idx, (llama_decision_order) t.decision_order);
}
@@ -2386,40 +2414,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() {
+17 -7
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
@@ -964,14 +964,19 @@ enum common_decision_type {
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
};
common_decision_type common_get_decision_type(const struct llama_model * model);
// same as above, but reads a GGUF file; it does not load the model
// returns COMMON_DECISION_TYPE_UNKNOWN if the file is missing, unreadable, or invalid
common_decision_type common_get_decision_type(const std::string & fname);
// metadata of a GGUF file, read without loading the model
struct common_gguf_info {
common_decision_type decision_type = COMMON_DECISION_TYPE_UNKNOWN; // UNKNOWN if the file is missing, unreadable, or invalid
uint32_t n_ctx_train = 0; // 0 if unknown
};
common_gguf_info common_get_gguf_info(const std::string & fname);
// note: defines the model, context, samplers, ets. lifetimes
struct common_init_result {
@@ -1069,6 +1074,7 @@ struct common_batch {
llama_seq_id seq_id; // the first sequence id, see add_seq()
bool output;
llama_embd embd; // non-owning view of the data passed to add_embd()/set_embd(), data == NULL if none
llama_embd state; // non-owning view of the data passed to set_embd_state(), data == NULL if none
std::vector<llama_seq_id> seq_ids_extra; // see add_seq()
int32_t decision_order = 0; // see llama_batch_ext_set_decision_order()
};
@@ -1106,6 +1112,9 @@ struct common_batch {
// attach a token embedding to the entry at idx, can only be set once per entry
bool set_embd(int32_t idx, llama_embd embd);
// attach a state embedding (e.g. the target hidden state for MTP) to the entry at idx, can only be set once per entry
bool set_embd_state(int32_t idx, llama_embd state);
// add an embedding-only entry (no token id)
// pos points to n_pos positions
int32_t add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output);
@@ -1295,12 +1304,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;
+2 -4
View File
@@ -77,7 +77,7 @@ common_chat_params common_chat_params_init_cohere2moe(const common_chat_template
data.prompt += data.generation_prompt;
}
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto generation_prompt = p.literal(GEN_PREFIX);
auto end = p.end();
@@ -124,12 +124,10 @@ common_chat_params common_chat_params_init_cohere2moe(const common_chat_template
return generation_prompt + reasoning + body + p.optional(p.literal(TURN_END)) + end;
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = !has_response_format && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
data.grammar_triggers = {
+2 -4
View File
@@ -145,7 +145,7 @@ common_chat_params common_chat_params_init_deepseek_v3_2(const common_chat_templ
bool require_tools = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED;
bool has_tool_calls = has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE;
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto generation_prompt = p.literal(GEN_PROMPT);
auto end = p.end();
@@ -256,12 +256,10 @@ common_chat_params common_chat_params_init_deepseek_v3_2(const common_chat_templ
generation_prompt + reasoning + content_before_tools + tool_calls + end;
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = has_tools && !require_tools;
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
data.grammar_triggers = {
+2 -4
View File
@@ -21,7 +21,7 @@ common_chat_params common_chat_params_init_functionary_v3_2(const common_chat_te
data.prompt += data.generation_prompt;
}
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
// Functionary v3.2 format:
// - Normal content: >>>all\n{content}
// - Tool calls: >>>function_name\n{json_args}
@@ -76,13 +76,11 @@ common_chat_params common_chat_params_init_functionary_v3_2(const common_chat_te
return generation_prompt + ret;
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
// Grammar trigger for when the model starts outputting a tool call
+2 -4
View File
@@ -198,7 +198,7 @@ common_chat_params common_chat_params_init_gemma4(const common_chat_template &
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto start = p.rule("start", p.optional(p.literal("<|turn>model\n")));
if (extract_reasoning) {
@@ -290,12 +290,10 @@ common_chat_params common_chat_params_init_gemma4(const common_chat_template &
return start + p.one_or_more(message);
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
data.grammar_triggers = {
+2 -4
View File
@@ -25,7 +25,7 @@ common_chat_params common_chat_params_init_gigachat_v3(
auto include_grammar = has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE;
const auto *tool_call_start_prefix = "<|message_sep|>\n\nfunction call<|role_sep|>\n";
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto ret = p.eps();
if (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE) {
// Build a choice of all available tools
@@ -60,13 +60,11 @@ common_chat_params common_chat_params_init_gigachat_v3(
return p.literal("assistant<|role_sep|>\n") + ret;
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
data.grammar_triggers = {
+3 -6
View File
@@ -45,8 +45,7 @@ common_chat_params common_chat_params_init_gpt_oss(const common_chat_template &
data.thinking_start_tag = "<|channel|>analysis<|message|>";
data.thinking_end_tags = {"<|end|>"};
// These special tokens are required to parse properly, so we include them
// even if parse_tool_calls is false.
// These special tokens are required to parse properly
data.preserved_tokens = {
"<|channel|>", "<|constrain|>", "<|message|>", "<|start|>", "<|end|>",
};
@@ -68,7 +67,7 @@ common_chat_params common_chat_params_init_gpt_oss(const common_chat_template &
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto start = p.rule("start", p.literal("<|start|>assistant"));
auto end = p.rule("end", p.literal("<|end|>"));
auto content = p.rule("message-content", p.until("<|end|>"));
@@ -138,12 +137,10 @@ common_chat_params common_chat_params_init_gpt_oss(const common_chat_template &
return p.zero_or_more(start + any) + start + (final_msg | unsolicited);
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
data.grammar_triggers = {
+2 -4
View File
@@ -77,7 +77,7 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template
data.prompt += data.generation_prompt;
}
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto generation_prompt = p.literal(GEN_PREFIX);
auto think_end = p.choice();
@@ -174,12 +174,10 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template
return generation_prompt + (reasoning << content << tool_calls);
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = !(has_response_format || inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED);
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
if (data.grammar_lazy) {
+2 -4
View File
@@ -48,7 +48,7 @@ common_chat_params common_chat_params_init_kimi_k2(const common_chat_template &
data.prompt += data.generation_prompt;
}
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
// Kimi K2 Thinking format:
// - Reasoning: <think>{reasoning}</think>
// - Content: text after reasoning
@@ -111,12 +111,10 @@ common_chat_params common_chat_params_init_kimi_k2(const common_chat_template &
return generation_prompt + reasoning + content_before_tools + tool_calls + end;
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
data.grammar_triggers = {
+2 -4
View File
@@ -66,7 +66,7 @@ common_chat_params common_chat_params_init_kimi_k3(const common_chat_template &
data.prompt += data.generation_prompt;
}
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto end = p.end();
auto start = p.optional(p.literal(MSG_START));
@@ -151,12 +151,10 @@ common_chat_params common_chat_params_init_kimi_k3(const common_chat_template &
return start + reasoning + response + tools + trailer + end;
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_REQUIRED;
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
data.grammar_triggers = {
+2 -4
View File
@@ -64,7 +64,7 @@ common_chat_params common_chat_params_init_lfm2(const common_chat_template &
data.prompt += data.generation_prompt;
}
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto generation_prompt = p.literal(GEN_PROMPT);
auto end = p.end();
@@ -93,12 +93,10 @@ common_chat_params common_chat_params_init_lfm2(const common_chat_template &
return generation_prompt + reasoning + content + tool_calls + end;
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
data.grammar_triggers = {
+2 -4
View File
@@ -80,7 +80,7 @@ common_chat_params common_chat_params_init_ling3(const common_chat_template &
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto end = p.end();
// the effective parse input is generation_prompt + model output, so the
@@ -185,12 +185,10 @@ common_chat_params common_chat_params_init_ling3(const common_chat_template &
return opener + reasoning + content + tools + tail + end;
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = !has_response_format && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_REQUIRED;
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
data.grammar_triggers = {
+3 -6
View File
@@ -48,8 +48,7 @@ common_chat_params common_chat_params_init_llm_jp_harmony(const common_chat_temp
data.thinking_start_tag = "<|channel|>analysis<|message|>";
data.thinking_end_tags = {"<|end|>"};
// These special tokens are required to parse properly, so we include them
// even if parse_tool_calls is false.
// These special tokens are required to parse properly
data.preserved_tokens = {
"<|channel|>", "<|constrain|>", "<|message|>", "<|start|>", "<|end|>",
};
@@ -71,7 +70,7 @@ common_chat_params common_chat_params_init_llm_jp_harmony(const common_chat_temp
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
// tokenizer space after special tokens; not p.space() since GBNF `space` allows one space only
auto sp = p.chars("[ ]", 0, -1);
auto channel_tag = p.literal("<|channel|>") + sp;
@@ -144,12 +143,10 @@ common_chat_params common_chat_params_init_llm_jp_harmony(const common_chat_temp
return p.zero_or_more(start + any) + start + final_msg;
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
data.grammar_triggers = {
+2 -4
View File
@@ -46,7 +46,7 @@ common_chat_params common_chat_params_init_minicpm5(const common_chat_template &
data.prompt += data.generation_prompt;
}
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto generation_prompt = p.literal("<|im_start|>assistant\n");
auto reasoning = p.eps();
@@ -113,12 +113,10 @@ common_chat_params common_chat_params_init_minicpm5(const common_chat_template &
return generation_prompt + reasoning + p.content(p.rest()) + p.end();
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
data.grammar_triggers = {
+2 -4
View File
@@ -56,7 +56,7 @@ common_chat_params common_chat_params_init_minimax_m3(const common_chat_template
data.prompt += data.generation_prompt;
}
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto generation_prompt = p.prefix(GEN_PROMPT, THINK_START);
auto end = p.end();
@@ -213,12 +213,10 @@ common_chat_params common_chat_params_init_minimax_m3(const common_chat_template
return generation_prompt + reasoning + content_before_tools + tool_calls + end;
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
data.grammar_triggers = {
+2 -4
View File
@@ -72,7 +72,7 @@ common_chat_params common_chat_params_init_ministral_3(const common_chat_templat
data.prompt += data.generation_prompt;
}
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto generation_prompt = p.eps();
auto reasoning =
extract_reasoning ? p.optional("[THINK]" + p.reasoning(p.until("[/THINK]")) + "[/THINK]") : p.eps();
@@ -108,13 +108,11 @@ common_chat_params common_chat_params_init_ministral_3(const common_chat_templat
return generation_prompt + (reasoning << p.content(p.rest()));
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
data.grammar_triggers = {
+2 -4
View File
@@ -48,7 +48,7 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
// Constrained grammar whenever tools are offered or a response format is requested.
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto start = p.rule("start", p.literal("<|start|>assistant"));
if (!extract_reasoning && !include_grammar) {
@@ -131,12 +131,10 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
return p.zero_or_more(start + analysis) + start + final_msg;
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
data.grammar_triggers = {
{ COMMON_GRAMMAR_TRIGGER_TYPE_PATTERN,
+2 -4
View File
@@ -71,7 +71,7 @@ common_chat_params common_chat_params_init_qwen3_coder(const common_chat_templat
});
}
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
auto generation_prompt = p.literal(GEN_PREFIX);
auto reasoning = p.eps();
@@ -174,13 +174,11 @@ common_chat_params common_chat_params_init_qwen3_coder(const common_chat_templat
return generation_prompt + (reasoning << p.content(p.rest()));
});
data.parser = parser.save();
if (include_grammar) {
data.grammar_lazy = has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
parser.build_grammar(builder, data.grammar_lazy);
data.parser.build_grammar(builder, data.grammar_lazy);
});
if (data.grammar_lazy) {
+1 -2
View File
@@ -54,10 +54,9 @@ common_chat_params common_chat_params_init_translate_gemma(
data.prompt += data.generation_prompt;
}
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
data.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;
}
-303
View File
@@ -1814,309 +1814,6 @@ void common_peg_arena::build_grammar(const common_grammar_builder & builder, boo
}
}
static common_json serialize_parser_variant(const common_peg_parser_variant & variant) {
using json = common_json;
return std::visit([](const auto & p) -> json {
using T = std::decay_t<decltype(p)>;
if constexpr (std::is_same_v<T, common_peg_epsilon_parser>) {
return json{{"type", "epsilon"}};
} else if constexpr (std::is_same_v<T, common_peg_start_parser>) {
return json{{"type", "start"}};
} else if constexpr (std::is_same_v<T, common_peg_end_parser>) {
return json{{"type", "end"}};
} else if constexpr (std::is_same_v<T, common_peg_literal_parser>) {
return json{{"type", "literal"}, {"literal", p.literal}};
} else if constexpr (std::is_same_v<T, common_peg_sequence_parser>) {
return json{{"type", "sequence"}, {"children", p.children}};
} else if constexpr (std::is_same_v<T, common_peg_choice_parser>) {
return json{{"type", "choice"}, {"children", p.children}};
} else if constexpr (std::is_same_v<T, common_peg_repetition_parser>) {
return json{
{"type", "repetition"},
{"child", p.child},
{"min_count", p.min_count},
{"max_count", p.max_count}
};
} else if constexpr (std::is_same_v<T, common_peg_and_parser>) {
return json{{"type", "and"}, {"child", p.child}};
} else if constexpr (std::is_same_v<T, common_peg_not_parser>) {
return json{{"type", "not"}, {"child", p.child}};
} else if constexpr (std::is_same_v<T, common_peg_any_parser>) {
return json{{"type", "any"}};
} else if constexpr (std::is_same_v<T, common_peg_space_parser>) {
return json{{"type", "space"}};
} else if constexpr (std::is_same_v<T, common_peg_chars_parser>) {
json ranges = json::array();
for (const auto & range : p.ranges) {
ranges.push_back({{"start", range.start}, {"end", range.end}});
}
return json{
{"type", "chars"},
{"pattern", p.pattern},
{"ranges", ranges},
{"negated", p.negated},
{"min_count", p.min_count},
{"max_count", p.max_count}
};
} else if constexpr (std::is_same_v<T, common_peg_string_parser>) {
return json{{"type", "string"}, {"delimiter", std::string(1, p.delimiter)}};
} else if constexpr (std::is_same_v<T, common_peg_until_parser>) {
return json{{"type", "until"}, {"delimiters", p.delimiters}};
} else if constexpr (std::is_same_v<T, common_peg_schema_parser>) {
return json{
{"type", "schema"},
{"child", p.child},
{"name", p.name},
{"raw", p.raw}
};
} else if constexpr (std::is_same_v<T, common_peg_rule_parser>) {
return json{
{"type", "rule"},
{"name", p.name},
{"child", p.child},
{"trigger", p.trigger}
};
} else if constexpr (std::is_same_v<T, common_peg_ref_parser>) {
return json{{"type", "ref"}, {"name", p.name}};
} else if constexpr (std::is_same_v<T, common_peg_atomic_parser>) {
return json{{"type", "atomic"}, {"child", p.child}};
} else if constexpr (std::is_same_v<T, common_peg_tag_parser>) {
return json{
{"type", "tag"},
{"child", p.child},
{"tag", p.tag}
};
} else if constexpr (std::is_same_v<T, common_peg_gbnf_parser>) {
return json{{"type", "gbnf"}, {"child", p.child}, {"grammar", p.grammar}};
} else if constexpr (std::is_same_v<T, common_peg_ac_parser>) {
return json{{"type", "ac"}, {"child", p.child}, {"delimiters", p.delimiters}};
}
}, variant);
}
common_json common_peg_arena::to_json() const {
auto parsers = common_json::array();
for (const auto & parser : parsers_) {
parsers.push_back(serialize_parser_variant(parser));
}
return common_json{
{"parsers", parsers},
{"rules", rules_},
{"root", root_}
};
}
static common_peg_parser_variant deserialize_parser_variant(const common_json & j) {
if (!j.contains("type") || !j["type"].is_string()) {
throw std::runtime_error("Parser variant JSON missing or invalid 'type' field");
}
std::string type = j["type"];
if (type == "epsilon") {
return common_peg_epsilon_parser{};
}
if (type == "start") {
return common_peg_start_parser{};
}
if (type == "end") {
return common_peg_end_parser{};
}
if (type == "literal") {
if (!j.contains("literal") || !j["literal"].is_string()) {
throw std::runtime_error("literal parser missing or invalid 'literal' field");
}
return common_peg_literal_parser{j["literal"]};
}
if (type == "sequence") {
if (!j.contains("children") || !j["children"].is_array()) {
throw std::runtime_error("sequence parser missing or invalid 'children' field");
}
return common_peg_sequence_parser{j["children"].get<std::vector<common_peg_parser_id>>()};
}
if (type == "choice") {
if (!j.contains("children") || !j["children"].is_array()) {
throw std::runtime_error("choice parser missing or invalid 'children' field");
}
return common_peg_choice_parser{j["children"].get<std::vector<common_peg_parser_id>>()};
}
if (type == "repetition") {
if (!j.contains("child") || !j.contains("min_count") || !j.contains("max_count")) {
throw std::runtime_error("repetition parser missing required fields");
}
return common_peg_repetition_parser{
j["child"].get<common_peg_parser_id>(),
j["min_count"].get<int>(),
j["max_count"].get<int>()
};
}
if (type == "and") {
if (!j.contains("child")) {
throw std::runtime_error("and parser missing 'child' field");
}
return common_peg_and_parser{j["child"].get<common_peg_parser_id>()};
}
if (type == "not") {
if (!j.contains("child")) {
throw std::runtime_error("not parser missing 'child' field");
}
return common_peg_not_parser{j["child"].get<common_peg_parser_id>()};
}
if (type == "any") {
return common_peg_any_parser{};
}
if (type == "space") {
return common_peg_space_parser{};
}
if (type == "chars") {
if (!j.contains("pattern") || !j.contains("ranges") || !j.contains("negated") ||
!j.contains("min_count") || !j.contains("max_count")) {
throw std::runtime_error("chars parser missing required fields");
}
common_peg_chars_parser parser;
parser.pattern = j["pattern"];
parser.negated = j["negated"].get<bool>();
parser.min_count = j["min_count"].get<int>();
parser.max_count = j["max_count"].get<int>();
for (const auto & range_json : j["ranges"]) {
if (!range_json.contains("start") || !range_json.contains("end")) {
throw std::runtime_error("char_range missing 'start' or 'end' field");
}
parser.ranges.push_back({
range_json["start"].get<uint32_t>(),
range_json["end"].get<uint32_t>()
});
}
return parser;
}
if (type == "string") {
if (!j.contains("delimiter")) {
throw std::runtime_error("string parser missing delimiter field.");
}
std::string delimiter = j["delimiter"];
if (delimiter.empty()) {
throw std::runtime_error("string parser delimiter is empty.");
}
return common_peg_string_parser{delimiter[0]};
}
if (type == "until") {
if (!j.contains("delimiters") || !j["delimiters"].is_array()) {
throw std::runtime_error("until parser missing or invalid 'delimiters' field");
}
return common_peg_until_parser{j["delimiters"].get<std::vector<std::string>>()};
}
if (type == "schema") {
if (!j.contains("child") || !j.contains("name") || !j.contains("raw")) {
throw std::runtime_error("schema parser missing required fields");
}
common_peg_schema_parser parser;
parser.child = j["child"].get<common_peg_parser_id>();
parser.name = j["name"];
parser.raw = j["raw"].get<bool>();
return parser;
}
if (type == "rule") {
if (!j.contains("name") || !j.contains("child") || !j.contains("trigger")) {
throw std::runtime_error("rule parser missing required fields");
}
return common_peg_rule_parser{
j["name"].get<std::string>(),
j["child"].get<common_peg_parser_id>(),
j["trigger"].get<bool>()
};
}
if (type == "ref") {
if (!j.contains("name") || !j["name"].is_string()) {
throw std::runtime_error("ref parser missing or invalid 'name' field");
}
return common_peg_ref_parser{j["name"]};
}
if (type == "atomic") {
if (!j.contains("child")) {
throw std::runtime_error("tag parser missing required fields");
}
return common_peg_atomic_parser{
j["child"].get<common_peg_parser_id>(),
};
}
if (type == "tag") {
if (!j.contains("child") || !j.contains("tag")) {
throw std::runtime_error("tag parser missing required fields");
}
return common_peg_tag_parser{
j["child"].get<common_peg_parser_id>(),
j["tag"].get<std::string>(),
};
}
if (type == "gbnf") {
if (!j.contains("child") || !j.contains("grammar")) {
throw std::runtime_error("gbnf parser missing required fields");
}
return common_peg_gbnf_parser{
j["child"].get<common_peg_parser_id>(),
j["grammar"].get<std::string>(),
};
}
if (type == "ac") {
if (!j.contains("child") || !j.contains("delimiters") || !j["delimiters"].is_array() || j["delimiters"].empty()) {
throw std::runtime_error("ac parser requires 'child' and a non-empty 'delimiters' array");
}
return common_peg_ac_parser{
j["child"].get<common_peg_parser_id>(),
j["delimiters"].get<std::vector<std::string>>(),
};
}
throw std::runtime_error("Unknown parser type: " + type);
}
common_peg_arena common_peg_arena::from_json(const common_json & j) {
if (!j.contains("parsers") || !j["parsers"].is_array()) {
throw std::runtime_error("JSON missing or invalid 'parsers' array");
}
if (!j.contains("rules") || !j["rules"].is_object()) {
throw std::runtime_error("JSON missing or invalid 'rules' object");
}
if (!j.contains("root")) {
throw std::runtime_error("JSON missing 'root' field");
}
common_peg_arena arena;
const auto & parsers_json = j["parsers"];
arena.parsers_.reserve(parsers_json.size());
for (const auto & parser_json : parsers_json) {
arena.parsers_.push_back(deserialize_parser_variant(parser_json));
}
arena.rules_ = j["rules"].get<std::unordered_map<std::string, common_peg_parser_id>>();
for (const auto & [name, id] : arena.rules_) {
if (id >= arena.parsers_.size()) {
throw std::runtime_error("Rule '" + name + "' references invalid parser ID: " + std::to_string(id));
}
}
arena.root_ = j["root"].get<common_peg_parser_id>();
if (arena.root_ != COMMON_PEG_INVALID_PARSER_ID && arena.root_ >= arena.parsers_.size()) {
throw std::runtime_error("Root references invalid parser ID: " + std::to_string(arena.root_));
}
return arena;
}
std::string common_peg_arena::save() const {
return to_json().dump();
}
void common_peg_arena::load(const std::string & data) {
*this = from_json(common_json::parse(data));
}
common_peg_arena build_peg_parser(const std::function<common_peg_parser(common_peg_parser_builder & builder)> & fn) {
common_peg_parser_builder builder;
builder.set_root(fn(builder));
-6
View File
@@ -357,12 +357,6 @@ class common_peg_arena {
std::string dump(common_peg_parser_id id) const;
common_json to_json() const;
static common_peg_arena from_json(const common_json & j);
std::string save() const;
void load(const std::string & data);
friend class common_peg_parser_builder;
private:
+38
View File
@@ -1050,3 +1050,41 @@ std::vector<common_sampler_type> common_sampler_types_from_chars(const std::stri
return samplers;
}
void common_sampling_add_preserved_tokens(common_params_sampling & sampling, const llama_vocab * vocab, const std::vector<std::string> & tokens) {
GGML_ASSERT(vocab != nullptr);
for (const auto & t : tokens) {
auto ids = common_tokenize(vocab, t, false, true);
if (ids.size() == 1) {
sampling.preserved_tokens.insert(ids[0]);
}
}
}
void common_sampling_add_grammar_triggers(common_params_sampling & sampling, const llama_vocab * vocab, std::vector<common_grammar_trigger> triggers) {
GGML_ASSERT(vocab != nullptr);
for (auto & trigger : triggers) {
if (trigger.type == COMMON_GRAMMAR_TRIGGER_TYPE_WORD) {
const auto & word = trigger.value;
auto ids = common_tokenize(vocab, word, false, true);
if (ids.size() == 1) {
auto token = ids[0];
if (std::find(sampling.preserved_tokens.begin(), sampling.preserved_tokens.end(), (llama_token) token) == sampling.preserved_tokens.end()) {
throw std::runtime_error("Grammar trigger word should be marked as preserved token: " + word);
}
common_grammar_trigger token_trigger;
token_trigger.type = COMMON_GRAMMAR_TRIGGER_TYPE_TOKEN;
token_trigger.value = word;
token_trigger.token = token;
sampling.grammar_triggers.push_back(std::move(token_trigger));
} else {
sampling.grammar_triggers.push_back({COMMON_GRAMMAR_TRIGGER_TYPE_WORD, word});
}
} else {
sampling.grammar_triggers.push_back(std::move(trigger));
}
}
if (sampling.grammar_lazy && sampling.grammar_triggers.empty()) {
throw std::runtime_error("Error: no triggers set for lazy grammar!");
}
}
+6
View File
@@ -118,6 +118,12 @@ std::string common_sampler_type_to_str(enum common_sampler_type cnstr);
std::vector<enum common_sampler_type> common_sampler_types_from_names(const std::vector<std::string> & names);
std::vector<enum common_sampler_type> common_sampler_types_from_chars(const std::string & chars);
// add the strings that are a single token in the vocab to the preserved tokens
void common_sampling_add_preserved_tokens(common_params_sampling & sampling, const llama_vocab * vocab, const std::vector<std::string> & tokens);
// add grammar triggers, a trigger word that is a single token becomes a token trigger and must be a preserved token
void common_sampling_add_grammar_triggers(common_params_sampling & sampling, const llama_vocab * vocab, std::vector<common_grammar_trigger> triggers);
llama_sampler * llama_sampler_init_llg(const llama_vocab * vocab,
const char * grammar_kind, const char * grammar_data);
+13 -9
View File
@@ -1541,8 +1541,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
return true;
}
// TODO: how to make it work with vision tokens?
if (!batch_in.has_token() || batch_in.has_embd()) {
if (!batch_in.has_token() && !batch_in.has_embd()) {
return true;
}
@@ -1581,15 +1580,20 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
const float * h_tgt = llama_get_embeddings_nextn(ctx_tgt);
for (int k = 0; k < n_tokens; ++k) {
const llama_seq_id seq_id = batch_in.tokens[k].seq_id;
const auto & t = batch_in.tokens[k];
const int32_t idx = batch.add(batch_in.tokens[k].id, batch_in.tokens[k].pos[0], seq_id, false);
const llama_seq_id seq_id = t.seq_id;
// vision tokens carry an embedding instead of an id
const int32_t idx = t.id != LLAMA_TOKEN_NULL
? batch.add(t.id, t.pos[0], seq_id, false)
: batch.add_embd(t.embd, t.pos.data(), seq_id, false);
const float * h_row = k == i_batch_beg[seq_id]
? pending_h[seq_id].data()
: h_tgt + (size_t) (k - 1) * n_embd;
batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
batch.set_embd_state(idx, { h_row, 1, (size_t) n_embd });
}
auto * mem_dft = llama_get_memory(ctx_dft);
@@ -1679,7 +1683,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
}
const int32_t idx = batch.add(dp.id_last, dp.pos0, seq_id, true);
batch.set_embd(idx, { pending_h[seq_id].data(), 1, (size_t) n_embd });
batch.set_embd_state(idx, { pending_h[seq_id].data(), 1, (size_t) n_embd });
i_last[seq_id] = idx;
@@ -1772,18 +1776,18 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
for (int t = 0; t < n_rows; ++t) {
const llama_token tok = (t == 0) ? dp.id_last : result[t - 1];
const int32_t idx = batch.add(tok, dp.pos0 + t, seq_id, t == n_rows - 1);
batch.set_embd(idx, { chain_h[seq_id].data() + (size_t) t * n_embd, 1, (size_t) n_embd });
batch.set_embd_state(idx, { chain_h[seq_id].data() + (size_t) t * n_embd, 1, (size_t) n_embd });
i_last[seq_id] = idx;
}
} else if (is_mem_shared) {
// note: with shared memory (e.g. Gemma4 assistants) we use the same position for all draft tokens
// ref: https://github.com/huggingface/transformers/blob/effde20942e3f82a1b97449f60b3a48c5ff96145/docs/source/en/model_doc/gemma4_assistant.md?plain=1#L36-L37
const int32_t idx = batch.add(id, dp.pos0, seq_id, true);
batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
batch.set_embd_state(idx, { h_row, 1, (size_t) n_embd });
i_last[seq_id] = idx;
} else {
const int32_t idx = batch.add(id, dp.pos0 + i + 1, seq_id, true);
batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
batch.set_embd_state(idx, { h_row, 1, (size_t) n_embd });
i_last[seq_id] = idx;
}
}
+5
View File
@@ -162,6 +162,7 @@ TEXT_MODEL_MAP: dict[str, str] = {
"Lfm2BidirectionalModel": "lfm2",
"Lfm2ForCausalLM": "lfm2",
"D1Model": "lfm2",
"D1OmniModel": "lfm2",
"Lfm2Model": "lfm2",
"Lfm2MoeForCausalLM": "lfm2",
"Llama4ForCausalLM": "llama",
@@ -188,6 +189,7 @@ TEXT_MODEL_MAP: dict[str, str] = {
"MiniCPM3ForCausalLM": "minicpm",
"MiniCPMForCausalLM": "minicpm",
"MiniCPMV4_6ForConditionalGeneration": "minicpm",
"MiniCPMV4_7ForConditionalGeneration": "minicpm",
"MiniMaxText01ForCausalLM": "minimax",
"MiniMaxM1ForCausalLM": "minimax",
"MiniMaxM2ForCausalLM": "minimax",
@@ -254,6 +256,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",
@@ -335,6 +338,7 @@ MMPROJ_MODEL_MAP: dict[str, str] = {
"KimiK25ForConditionalGeneration": "kimivl",
"KimiVLForConditionalGeneration": "kimivl",
"Lfm2AudioForConditionalGeneration": "lfm2",
"D1OmniModel": "lfm2",
"Lfm2VlForConditionalGeneration": "lfm2",
"LightOnOCRForConditionalGeneration": "lighton_ocr",
"Llama4ForConditionalGeneration": "llama4",
@@ -343,6 +347,7 @@ MMPROJ_MODEL_MAP: dict[str, str] = {
"MiMoV2ForCausalLM": "mimo",
"MiniMaxM3SparseForConditionalGeneration": "minimax",
"MiniCPMV4_6ForConditionalGeneration": "minicpm",
"MiniCPMV4_7ForConditionalGeneration": "minicpm",
"Mistral3ForConditionalGeneration": "llava",
"NemotronH_Nano_VL_V2": "nemotron",
"MuseGlimmerForConditionalGeneration": "muse_glimmer",
+61 -45
View File
@@ -439,6 +439,25 @@ class ModelBase:
return (unpacked * scale.unsqueeze(-1).float()).reshape(shape)
def dequant_fp8() -> None:
for name in self.model_tensors.keys():
if name.endswith(".weight_scale"):
weight_name = name.removesuffix("_scale")
if weight_name not in self.model_tensors:
tensors_to_remove.append(name)
continue
w = self.model_tensors[weight_name]
s = self.model_tensors[name]
is_fp8_weight = False
if self._fp8_as_q8:
is_fp8_weight = w().dtype in (torch.float8_e4m3fn, torch.float8_e5m2)
self.model_tensors[weight_name] = lambda w=w, s=s: dequant_simple(w(), s(), None)
tensors_to_remove.append(name)
if is_fp8_weight:
self._fp8_dequantized.add(weight_name)
if name.endswith((".input_scale", ".k_scale", ".v_scale")):
tensors_to_remove.append(name)
if quant_method == "bitnet":
for name in self.model_tensors.keys():
if name.endswith(".weight_scale"):
@@ -498,18 +517,14 @@ class ModelBase:
elif quant_method == "compressed-tensors":
quant_format = quant_config["format"]
groups = quant_config["config_groups"]
nvfp4_compressed_tensors = (
quant_format == "nvfp4-pack-quantized"
or quant_format == "mixed-precision"
and bool(groups)
and all(g.get("format") == "nvfp4-pack-quantized" for g in groups.values() if isinstance(g, dict))
)
nvfp4_compressed_tensors = self._is_nvfp4_compressed_tensors(quant_method, quant_format, groups)
if len(groups) > 1 and not nvfp4_compressed_tensors:
if nvfp4_compressed_tensors:
dequant_fp8()
elif len(groups) > 1:
raise NotImplementedError("Can't handle multiple config groups for compressed-tensors yet")
weight_config = tuple(groups.values())[0]["weights"]
if quant_format == "float-quantized" or quant_format == "int-quantized" or quant_format == "naive-quantized":
elif quant_format == "float-quantized" or quant_format == "int-quantized" or quant_format == "naive-quantized":
weight_config = tuple(groups.values())[0]["weights"]
block_size = weight_config.get("block_structure", None)
strategy = weight_config.get("strategy")
assert strategy == "channel" or strategy == "block"
@@ -529,6 +544,7 @@ class ModelBase:
if self._fp8_as_q8 and is_fp8:
self._fp8_dequantized.add(weight_name)
elif quant_format == "pack-quantized":
weight_config = tuple(groups.values())[0]["weights"]
assert weight_config.get("strategy") == "group"
assert weight_config.get("type", "int") == "int"
num_bits = weight_config.get("num_bits")
@@ -550,32 +566,10 @@ class ModelBase:
tensors_to_remove += [base_name + n for n in ("_packed", "_shape", "_scale")]
if (base_name + "_zero_point") in self.model_tensors:
tensors_to_remove.append(base_name + "_zero_point")
elif nvfp4_compressed_tensors:
# Don't error from compressed-tensors, we'll handle them in _generate_nvfp4_tensors
pass
else:
raise NotImplementedError(f"Quant format {quant_format!r} for method {quant_method!r} is not yet supported")
elif quant_method == "modelopt":
# Mixed-precision ModelOpt models: NVFP4 tensors are handled by
# _generate_nvfp4_tensors; FP8 tensors have 1D weight_scale and
# are dequantized here. k/v scale tensors are unused.
for name in self.model_tensors.keys():
if name.endswith(".weight_scale"):
weight_name = name.removesuffix("_scale")
if weight_name not in self.model_tensors:
tensors_to_remove.append(name)
continue
w = self.model_tensors[weight_name]
s = self.model_tensors[name]
is_fp8_weight = False
if self._fp8_as_q8:
is_fp8_weight = w().dtype in (torch.float8_e4m3fn, torch.float8_e5m2)
self.model_tensors[weight_name] = lambda w=w, s=s: dequant_simple(w(), s(), None)
tensors_to_remove.append(name)
if is_fp8_weight:
self._fp8_dequantized.add(weight_name)
if name.endswith((".input_scale", ".k_scale", ".v_scale")):
tensors_to_remove.append(name)
dequant_fp8()
elif quant_method is not None:
raise NotImplementedError(f"Quant method is not yet supported: {quant_method!r}")
@@ -821,6 +815,18 @@ class ModelBase:
func=load,
)
@staticmethod
def _is_nvfp4_compressed_tensors(quant_method, quant_format, groups) -> bool:
# Some models use per-tensor quant_algo (e.g. "MIXED_PRECISION" with
# per-layer NVFP4/FP8) instead of a single global "NVFP4" value.
if quant_method != "compressed-tensors":
return False
if quant_format == "nvfp4-pack-quantized":
return True
if quant_format != "mixed-precision" or not groups:
return False
return any(g.get("format") == "nvfp4-pack-quantized" for g in groups.values() if isinstance(g, dict))
@staticmethod
def _nvfp4_pack(weight: Tensor, scale: Tensor) -> tuple[np.ndarray, list[int]]:
"""Repack NVFP4 ModelOpt tensors into ggml super-block layout.
@@ -878,8 +884,8 @@ class ModelBase:
weight = LazyTorchTensor.to_eager(self.model_tensors[name]())
scale = LazyTorchTensor.to_eager(self.model_tensors[scale_name]())
# Skip non-NVFP4 tensors (e.g. FP8 with per-channel 1D scales)
if scale.ndim < 2:
# Skip non-NVFP4 tensors(e.g. 1D scale, or float8 weight)
if scale.ndim < 2 or weight.dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
continue
scale2 = LazyTorchTensor.to_eager(self.model_tensors.get(scale2_name, lambda: torch.tensor(1.0))())
@@ -980,14 +986,7 @@ class ModelBase:
quant_groups = quant_config.get("config_groups", quant_groups) or {}
quant_layers = quant_config.get("quantized_layers", quant_layers) or {}
# Some models use per-tensor quant_algo (e.g. "MIXED_PRECISION" with
# per-layer NVFP4/FP8) instead of a single global "NVFP4" value.
nvfp4_compressed_tensors = quant_method == "compressed-tensors" and (
quant_format == "nvfp4-pack-quantized"
or quant_format == "mixed-precision"
and bool(quant_groups)
and all(g.get("format") == "nvfp4-pack-quantized" for g in quant_groups.values() if isinstance(g, dict))
)
nvfp4_compressed_tensors = self._is_nvfp4_compressed_tensors(quant_method, quant_format, quant_groups)
self._nvfp4_global_algo = quant_algo
@@ -1390,14 +1389,17 @@ class TextModel(ModelBase):
name, gen = item
# Skip multimodal tensors
if name.startswith(("mlp", "vit.", "vpm.", "siglip2.", "conformer.", "merger.", "resampler.", "sound_encoder.", "sound_projection.", "speech_embeddings.")) \
# strip the "model." wrapper so the prefixes below match (name is not returned)
if name.startswith("model."):
name = name[len("model."):]
if name.startswith(("mlp", "vit.", "vpm.", "siglip2.", "conformer.", "connector.", "merger.", "resampler.", "sound_encoder.", "sound_projection.", "speech_embeddings.")) \
or "visual." in name or "vision." in name or "audio." in name or "talker." in name \
or "vision_" in name or "audio_" in name \
or "token2wav." in name or "code2wav." in name \
or "projector." in name or "pre_mm_projector_norm" in name \
or "image_newline" in name or "view_seperator" in name \
or "patch_embed" in name or "patch_embedding" in name \
or "patch_merger." in name or "patch_merge_mlp." in name or "model.connector." in name:
or "patch_merger." in name or "patch_merge_mlp." in name:
return None
return super().filter_tensors(item)
@@ -2336,12 +2338,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"))
+167
View File
@@ -161,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):
@@ -276,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):
+85 -4
View File
@@ -139,9 +139,16 @@ class MiniCPMV4_6TextModel(Qwen3_5TextModel):
@ModelBase.register("MiniCPMV4_6ForConditionalGeneration")
@ModelBase.example("openbmb/MiniCPM-V-4_6")
class MiniCPMV4_6VisionModel(MmprojModel):
projector_type = gguf.VisionProjectorType.MINICPMV4_6
# fallback for checkpoints whose preprocessor config omits `scale_resolution`
default_scale_resolution: int | None = None
def get_downsample_mode(self) -> str:
return self.preprocessor_config.get("downsample_mode", "16x")
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.downsample_mode = self.preprocessor_config.get("downsample_mode", "16x")
self.downsample_mode = self.get_downsample_mode()
if self.downsample_mode not in {"4x", "16x"}:
raise ValueError(f"Unsupported downsample mode: {self.downsample_mode}")
if self.downsample_mode == "4x":
@@ -157,7 +164,8 @@ class MiniCPMV4_6VisionModel(MmprojModel):
# The CLIP loader in tools/mtmd/clip.cpp consumes `clip.vision.image_size`
# as the slice size and warmup resolution, so report `scale_resolution` there
# to match the upstream MiniCPMV4_6ImageProcessorPil slicing rules.
scale_resolution = self.preprocessor_config.get("scale_resolution")
scale_resolution = self.preprocessor_config.get(
"scale_resolution", self.default_scale_resolution)
if scale_resolution is not None:
self.hparams_vision["image_size"] = int(scale_resolution)
@@ -166,12 +174,15 @@ class MiniCPMV4_6VisionModel(MmprojModel):
assert self.hparams_vision is not None
# projector type string is consumed by clip_projector_type_from_string() in clip.cpp
# (mapped to PROJECTOR_TYPE_MINICPMV4_6).
self.gguf_writer.add_clip_projector_type(gguf.VisionProjectorType.MINICPMV4_6)
self.gguf_writer.add_clip_projector_type(self.projector_type)
self.gguf_writer.add_vision_projector_scale_factor(
2 if self.downsample_mode == "4x" else 4)
max_slice_nums = self.preprocessor_config.get("max_slice_nums")
if max_slice_nums is not None:
self.gguf_writer.add_vision_max_slice_nums(int(max_slice_nums))
# borrow wa_layer_indexes for vit_merger insertion point
insert_layer_id = int(self.global_config.get(
"insert_layer_id", self.hparams_vision.get("insert_layer_id", 6)))
@@ -191,3 +202,73 @@ class MiniCPMV4_6VisionModel(MmprojModel):
return None
return super().filter_tensors(item)
# MiniCPM-V 4.7 shares the v4.6 stack: the same Qwen3.5 text tower (MoE variant when the checkpoint says so) and the same SigLIP + vit_merger + merger vision tower.
@ModelBase.register("MiniCPMV4_7ForConditionalGeneration")
@ModelBase.example("openbmb/MiniCPM-V-4.7")
class MiniCPMV4_7TextModel(Qwen3_5TextModel):
model_arch = gguf.MODEL_ARCH.QWEN35
def set_gguf_parameters(self):
super().set_gguf_parameters()
# mtmd puts the time of the image canvas in slot z, slot t stays the KV cache position
self.gguf_writer.add_rope_section_order(gguf.RopeSectionOrder.ZYXT)
def __init__(self, dir_model, ftype, fname_out, *, hparams: dict | None = None, **kwargs):
if hparams is None:
hparams = ModelBase.load_hparams(dir_model, is_mistral_format=False)
text_config = hparams.get("text_config", {})
if text_config.get("model_type") == "qwen3_5_moe_text":
self.model_arch = gguf.MODEL_ARCH.QWEN35MOE
else:
self.model_arch = gguf.MODEL_ARCH.QWEN35
super().__init__(dir_model, ftype, fname_out, hparams=hparams, **kwargs)
@classmethod
def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None:
name, gen = item
# MTP tensors are not used yet
if name.startswith("mtp"):
return None
return super().filter_tensors(item)
@ModelBase.register("MiniCPMV4_7ForConditionalGeneration")
@ModelBase.example("openbmb/MiniCPM-V-4.7")
class MiniCPMV4_7VisionModel(MiniCPMV4_6VisionModel):
projector_type = gguf.VisionProjectorType.MINICPMV4_7
# MiniCPMV4_7ImageProcessorPil default
default_scale_resolution = 448
# rows of v.tok_embd_sep, the order must match clip_suffix_rows() in clip-impl.h
tok_embd_sep = ["</image>", "<slice>", "</slice>", "\n"]
def get_downsample_mode(self) -> str:
# 4.7 moved downsample_mode to the model config; preprocessor value takes priority
return self.preprocessor_config.get(
"downsample_mode", self.global_config.get("downsample_mode", "16x"))
@classmethod
def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None:
# keep the text tok_embd, the separator rows are taken from it in modify_tensors
if item[0] == "model.language_model.embed_tokens.weight":
return item
return super().filter_tensors(item)
def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:
if name == "model.language_model.embed_tokens.weight":
# the tile separators are text tokens; clip appends their embeddings so that one chunk holds the whole image
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(self.dir_model)
ids = []
for text in self.tok_embd_sep:
tok = tokenizer.encode(text, add_special_tokens=False)
if len(tok) != 1:
raise ValueError(f"separator {text!r} must be a single token, got {tok}")
ids.append(tok[0])
yield self.format_tensor_name(gguf.MODEL_TENSOR.V_TOK_EMBD_SEP, suffix=""), data_torch[ids]
return
yield from super().modify_tensors(data_torch, name, bid)
+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"
+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."},
+55
View File
@@ -0,0 +1,55 @@
## MiniCPM-V 4.7
### Prepare models and code
Download [MiniCPM-V-4.7](https://huggingface.co/openbmb/MiniCPM-V-4.7) PyTorch model from huggingface to "MiniCPM-V-4.7" folder.
The model must be the standard `transformers` checkpoint (no `trust_remote_code` for the text and vision graph used here); the architecture in `config.json` is `MiniCPMV4_7ForConditionalGeneration` with a `qwen3_5_text` (or `qwen3_5_moe_text`) text model and a SigLIP-based vision tower plus a window-attention `vit_merger`, same as MiniCPM-V 4.6.
If the checkpoint ships no MTP weights, pass `--no-mtp` to skip the nextn layers.
### Build llama.cpp
If there are differences in usage, please refer to the official build [documentation](https://github.com/ggml-org/llama.cpp/blob/master/docs/build.md)
Clone llama.cpp:
```bash
git clone https://github.com/ggml-org/llama.cpp
cd llama.cpp
```
Build llama.cpp using `CMake`:
```bash
cmake -B build
cmake --build build --config Release
```
### Usage of MiniCPM-V 4.7
MiniCPM-V 4.7 is converted directly through `convert_hf_to_gguf.py`. The same script is invoked twice on the original Hugging Face directory: once to produce the language-model GGUF and once with `--mmproj` to produce the multimodal projector GGUF.
```bash
# language model
python ./convert_hf_to_gguf.py ../MiniCPM-V-4.7 --outfile ../MiniCPM-V-4.7/ggml-model-f16.gguf --no-mtp
# multimodal projector (vision tower + window-attention vit_merger + DownsampleMLP merger)
python ./convert_hf_to_gguf.py ../MiniCPM-V-4.7 --mmproj --outfile ../MiniCPM-V-4.7/mmproj-model-f16.gguf
# optional: quantize to Q4_K_M
./build/bin/llama-quantize ../MiniCPM-V-4.7/ggml-model-f16.gguf ../MiniCPM-V-4.7/ggml-model-Q4_K_M.gguf Q4_K_M
```
The default projector merges 16x (4x4 patches into one token). To keep 4x more visual tokens, copy the model dir and set `"downsample_mode": "4x"` in the copy's `preprocessor_config.json` before running the `--mmproj` conversion; the loader reads `clip.vision.projector.scale_factor` to pick the graph.
Inference on Linux or Mac
```bash
# run in single-turn mode
./build/bin/llama-mtmd-cli -m ../MiniCPM-V-4.7/ggml-model-f16.gguf --mmproj ../MiniCPM-V-4.7/mmproj-model-f16.gguf -c 4096 --jinja --image xx.jpg -p "What is in the image?"
# run in conversation mode
./build/bin/llama-mtmd-cli -m ../MiniCPM-V-4.7/ggml-model-Q4_K_M.gguf --mmproj ../MiniCPM-V-4.7/mmproj-model-f16.gguf --jinja
```
The chat template enables thinking by default. Pass `--chat-template-kwargs '{"enable_thinking": false}'` to `llama-server` to turn it off.
+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);
}
+10
View File
@@ -462,6 +462,8 @@ function(ggml_add_cpu_backend_variant tag_name)
set(GGML_INTERNAL_${feat} ON)
endforeach()
elseif (GGML_SYSTEM_ARCH STREQUAL "s390x")
set(GGML_NATIVE OFF)
foreach (feat VXE2 NNPA)
set(GGML_INTERNAL_${feat} OFF)
endforeach()
@@ -569,6 +571,14 @@ if (GGML_CPU_ALL_VARIANTS)
if (CMAKE_SYSTEM_NAME MATCHES "Linux")
ggml_add_cpu_backend_variant(z15 Z15 VXE2)
ggml_add_cpu_backend_variant(z16 Z16 VXE2 NNPA)
# check if compiler supports "-march=z17" codename
check_cxx_compiler_flag("-march=arch15" GGML_CXX_SUPPORTS_Z17)
if (GGML_CXX_SUPPORTS_Z17)
ggml_add_cpu_backend_variant(arch15 Z17 VXE2 NNPA)
else()
message(WARNING "Skipping z17 target: compiler must be GCC 15.1 and later")
endif()
else()
message(FATAL_ERROR "Unsupported s390x target OS: ${CMAKE_SYSTEM_NAME}")
endif()
+20 -3
View File
@@ -1176,7 +1176,22 @@ static struct ggml_backend_meta_split_state ggml_backend_meta_get_split_state(
return ret;
}
static bool ggml_backend_meta_is_host_view(const struct ggml_tensor * tensor) {
return ggml_is_view(tensor) && ggml_backend_buffer_is_host(tensor->view_src->buffer);
}
static struct ggml_backend_meta_split_state ggml_backend_meta_get_split_state(const struct ggml_tensor * tensor, bool assume_sync) {
// [TAG_META_HOST_VIEWS]
// TODO: technically, this check should not be needed if the backend scheduler correctly prevents assigning
// such host-buffer views to the meta backend. figure out how to update the scheduler logic to achieve that
// ref: https://github.com/ggml-org/llama.cpp/pull/30217
if (!ggml_backend_buffer_is_meta(tensor->buffer)) {
GGML_ASSERT(ggml_backend_meta_is_host_view(tensor));
// the view is not allocated in the meta buffer, it is not split across the sub-devices
return { GGML_BACKEND_SPLIT_AXIS_NONE, {0}, {1}, 1 };
}
ggml_backend_meta_buffer_context * buf_ctx = (ggml_backend_meta_buffer_context *) tensor->buffer->context;
return ggml_backend_meta_get_split_state(buf_ctx->get_simple_tensor_container(tensor), tensor, assume_sync);
}
@@ -2026,9 +2041,11 @@ static enum ggml_status ggml_backend_meta_graph_compute(ggml_backend_t backend,
for (int i = 0; i < cgraph->n_nodes; i++) {
ggml_tensor * node = cgraph->nodes[i];
if (node->view_src != nullptr && node->view_src->op == GGML_OP_NONE && ggml_backend_buffer_is_host(node->view_src->buffer)) {
// FIXME s_copy_main is on the CPU and its view seems to be incorrectly added to the graph nodes.
// For regular usage this doesn't matter since it's a noop but trying to call ggml_backend_meta_buffer_simple_tensor results in a crash.
if (!ggml_backend_buffer_is_meta(node->buffer)) {
// [TAG_META_HOST_VIEWS]
GGML_ASSERT(ggml_backend_meta_is_host_view(node));
// keep the node as is, mapping it to a simple tensor is not possible
bcj.nodes[i] = node;
continue;
}
+6 -1
View File
@@ -593,7 +593,12 @@ function(ggml_add_cpu_backend_variant_impl tag_name)
foreach (ZHW RANGE 15 17)
if(DEFINED GGML_INTERNAL_Z${ZHW})
message(STATUS "z${ZHW} cross-compile target")
list(APPEND ARCH_FLAGS -march=z${ZHW})
if (ZHW EQUAL 17)
# z17 is an alias of arch15, use the arch level for wider toolchain support
list(APPEND ARCH_FLAGS -march=arch15)
else()
list(APPEND ARCH_FLAGS -march=z${ZHW})
endif()
endif()
endforeach()
endif()
+2 -3
View File
@@ -390,17 +390,16 @@ typedef unsigned char uchar8x16_t __attribute__((vector_size(16)));
typedef int8_t int8x16_t __attribute__((vector_size(16)));
typedef int16_t int16x8_t __attribute__((vector_size(16)));
typedef int32_t int32x4_t __attribute__((vector_size(16)));
typedef int64_t int64x2_t __attribute__((vector_size(16)));
typedef uint8_t uint8x16_t __attribute__((vector_size(16)));
typedef uint16_t uint16x8_t __attribute__((vector_size(16)));
typedef uint32_t uint32x4_t __attribute__((vector_size(16)));
typedef uint64_t uint64x2_t __attribute__((vector_size(16)));
typedef float float32x4_t __attribute__((vector_size(16)));
typedef double double64x2_t __attribute__((vector_size(16)));
typedef signed long long long64x2_t __attribute__((vector_size(16)));
typedef unsigned long long ulong64x2_t __attribute__((vector_size(16)));
typedef struct ggml_uint8x16x2_t {
uint8x16_t val[2];
} ggml_uint8x16x2_t;
+6
View File
@@ -3504,6 +3504,12 @@ void ggml_cpu_fp32_to_fp16(const float * x, ggml_fp16_t * y, int64_t n) {
vfloat16m1_t vy = __riscv_vfncvt_f_f_w_f16m1(vx, vl);
__riscv_vse16_v_f16m1((_Float16 *)&y[i], vy, vl);
}
#elif defined(__VXE__) || defined(__VXE2__)
for (; i + 7 < n; i += 8) {
const uint32x4_t v_yl = __lzs_f32cx4_to_f16(vec_xl(0, x + i + 0));
const uint32x4_t v_yh = __lzs_f32cx4_to_f16(vec_xl(0, x + i + 4));
vec_xst(vec_pack(v_yl, v_yh), 0, (uint16_t *)(y + i));
}
#endif
for (; i < n; ++i) {
y[i] = GGML_CPU_FP32_TO_FP16(x[i]);
+21 -9
View File
@@ -1223,6 +1223,24 @@ static inline void __lsx_f16x4_store(ggml_fp16_t * x, __m128 y) {
#define GGML_F16_STEP GGML_F32_STEP
#define GGML_F16_EPR GGML_F32_EPR
static inline uint32x4_t __lzs_f32cx4_to_f16(float32x4_t v_f) {
float32x4_t v_base = vec_mul(vec_mul(vec_abs(v_f), vec_splats(0x1.0p+112f)), vec_splats(0x1.0p-110f));
const uint32x4_t v_w = (uint32x4_t)v_f;
const uint32x4_t v_shl1_w = vec_add(v_w, v_w);
const uint32x4_t v_sign = vec_and(v_w, vec_splats(UINT32_C(0x80000000)));
const uint32x4_t v_bias = vec_max(vec_and(v_shl1_w, vec_splats(UINT32_C(0xFF000000))), vec_splats(UINT32_C(0x71000000)));
v_base = vec_add((float32x4_t)vec_add(vec_sr(v_bias, 1), vec_splats(UINT32_C(0x07800000))), v_base);
const uint32x4_t v_bits = (uint32x4_t)v_base;
const uint32x4_t v_nonsign = vec_add(vec_and(vec_sr(v_bits, 13), vec_splats(UINT32_C(0x00007C00))),
vec_and(v_bits, vec_splats(UINT32_C(0x00000FFF))));
const uint32x4_t v_is_nan = (uint32x4_t)vec_cmpgt(v_shl1_w, vec_splats(UINT32_C(0xFF000000)));
return vec_or(vec_sr(v_sign, 16), vec_sel(v_nonsign, vec_splats(UINT32_C(0x7E00)), v_is_nan));
}
static inline float32x4_t __lzs_f16cx4_load(const ggml_fp16_t * x) {
float tmp[4];
@@ -1236,15 +1254,9 @@ static inline float32x4_t __lzs_f16cx4_load(const ggml_fp16_t * x) {
}
static inline void __lzs_f16cx4_store(ggml_fp16_t * x, float32x4_t v_y) {
float arr[4];
// note: keep type-cast here to prevent compiler bugs
// see: https://github.com/ggml-org/llama.cpp/issues/12846
vec_xst(v_y, 0, (float *)(arr));
for (int i = 0; i < 4; i++) {
x[i] = GGML_CPU_FP32_TO_FP16(arr[i]);
}
const uint32x4_t v_h = __lzs_f32cx4_to_f16(v_y);
const uint64_t tmp = ((uint64x2_t)vec_pack(v_h, v_h))[0];
memcpy(x, &tmp, sizeof(tmp));
}
#define GGML_F16_VEC GGML_F32x4
+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,
+1 -1
View File
@@ -1772,7 +1772,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile(
}
}
}
if (np > 1) {
if (np > 1 || nbatch_combine != DKQ/2) {
__syncthreads();
}
}
+3
View File
@@ -200,9 +200,12 @@ static bool ggml_cuda_op_fwht_impl(ggml_backend_cuda_context & ctx, const ggml_t
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);
+94 -10
View File
@@ -2899,6 +2899,79 @@ static int ggml_cuda_try_gdn_cache_fusion(
return skip;
}
// match ssm_scan + the strided cpy that scatters its state snapshots into the cache, so the kernel writes them and skips the cpy
static int ggml_cuda_try_ssm_scan_cache_fusion(
const ggml_cgraph * cgraph, int node_idx, ggml_cuda_ssm_scan_fused_cache & fused_state_cpy) {
const ggml_tensor * ssm = cgraph->nodes[node_idx];
// the kernel skips the snapshot tail, so the scan output must not be a graph output
if (ssm->op != GGML_OP_SSM_SCAN || ssm->type != GGML_TYPE_F32 || (ssm->flags & GGML_TENSOR_FLAG_OUTPUT)) {
return 0;
}
const int64_t K = ggml_get_op_params_i32(ssm, 0); // snapshot slot count
const ggml_tensor * s = ssm->src[0];
const ggml_tensor * x = ssm->src[1];
const ggml_tensor * A = ssm->src[3];
const int64_t d_state = s->ne[0];
const int64_t D = d_state * s->ne[1] * x->ne[1]; // d_state * head_dim * n_head
const int64_t n_tok = x->ne[2];
const int64_t n_seqs = x->ne[3];
// only the mamba-2 kernels (group scan and SSD) write to the cache; mamba-1 still uses the cpy
if (A->nb[1] != sizeof(float) || (d_state != 96 && d_state != 128 && d_state != 256)) {
return 0;
}
// the scan reads its input rows from the cache (picked by ids), so with more than one seq a seq can read a row that another seq writes in the same launch
if (n_seqs != 1) {
return 0;
}
const int64_t n_written = std::min<int64_t>(n_tok, K);
const size_t tail_off = ggml_row_size(GGML_TYPE_F32, ggml_nelements(x));
// snapshot cpy is the first real node after the scan (skip views/no-ops)
const ggml_tensor * cpy = nullptr;
int skip = 0;
for (int j = node_idx + 1; j < cgraph->n_nodes && cpy == nullptr; ++j) {
const ggml_tensor * n = cgraph->nodes[j];
if (ggml_cuda_is_view_or_noop(n)) {
continue;
}
if (n->op != GGML_OP_CPY || (n->flags & GGML_TENSOR_FLAG_OUTPUT)) {
return 0;
}
cpy = n;
skip = j - node_idx;
}
if (cpy == nullptr) {
return 0;
}
const ggml_tensor * src = cpy->src[0]; // view of the scan snapshot tail
const ggml_tensor * dst = cpy->src[1]; // cache view the kernel writes to
// src must be this scan's snapshot tail (contiguous, at the tail offset)
if (src->op != GGML_OP_VIEW || src->view_src != ssm || src->view_offs != tail_off ||
!ggml_is_contiguous(src)) {
return 0;
}
// dst is the [D, n_seqs, n_written] cache view; require nb[1] == D, the per-seq stride the kernel takes from src0->nb[3]
const std::array<int64_t, GGML_MAX_DIMS> expected_ne = { D, n_seqs, n_written, 1 };
if (dst->op != GGML_OP_VIEW || dst->type != GGML_TYPE_F32 || dst->data == nullptr ||
!std::equal(expected_ne.begin(), expected_ne.end(), dst->ne) ||
dst->nb[0] != ggml_type_size(GGML_TYPE_F32) || dst->nb[1] != (size_t) ggml_row_size(GGML_TYPE_F32, D)) {
return 0;
}
fused_state_cpy.data = (float *) dst->data; // rollback slot 0 (newest)
fused_state_cpy.slot_stride = K > 1 ? (int64_t) (dst->nb[2] / sizeof(float)) : 0;
return skip;
}
static bool ggml_cuda_topk_moe_fusion(const struct ggml_cgraph * cgraph, int node_idx, ggml_cuda_topk_moe_args & args) {
args.sigmoid = false;
args.sqrt_softplus = false;
@@ -3585,6 +3658,20 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph
}
}
// ssm_scan -> cpy: scatter recurrent-state snapshots into the cache
if (node->op == GGML_OP_SSM_SCAN) {
ggml_cuda_ssm_scan_fused_cache fused_state_cpy;
const int nodes_to_skip = ggml_cuda_try_ssm_scan_cache_fusion(cgraph, i, fused_state_cpy);
if (nodes_to_skip > 0) {
#ifdef GGML_CUDA_DEBUG
GGML_LOG_INFO("%s: fused ssm_scan snapshot copies for %s (skipped %d nodes)\n",
__func__, node->name, nodes_to_skip);
#endif
ggml_cuda_op_ssm_scan_fused_cache(*cuda_ctx, node, fused_state_cpy);
return nodes_to_skip;
}
}
//topk-moe
if (cgraph->nodes[i]->op == GGML_OP_UNARY || cgraph->nodes[i]->op == GGML_OP_SOFT_MAX ||
cgraph->nodes[i]->op == GGML_OP_ARGSORT) {
@@ -5314,9 +5401,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 +5746,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 +5776,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 +5788,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) \
+3 -3
View File
@@ -283,9 +283,9 @@ static constexpr __host__ __device__ int get_mmvq_mmid_max_batch_rdna4(ggml_type
// Host function: returns the max batch size for the current arch+type at runtime.
int get_mmvq_mmid_max_batch(ggml_type type, int cc) {
// NVIDIA: Volta, Ada Lovelace, and Blackwell always use MMVQ for MUL_MAT_ID.
// NVIDIA: P100, Volta, Ada Lovelace, and Blackwell always use MMVQ for MUL_MAT_ID.
if (GGML_CUDA_CC_IS_NVIDIA(cc)) {
if (cc == GGML_CUDA_CC_VOLTA || cc >= GGML_CUDA_CC_ADA_LOVELACE) {
if (cc == GGML_CUDA_CC_PASCAL || cc == GGML_CUDA_CC_VOLTA || cc >= GGML_CUDA_CC_ADA_LOVELACE) {
return MMVQ_MAX_BATCH_SIZE;
}
if (cc >= GGML_CUDA_CC_TURING) {
@@ -440,7 +440,7 @@ static constexpr __device__ int get_mmvq_mmid_max_batch_for_device() {
return get_mmvq_mmid_max_batch_cdna(type);
#elif defined(GCN)
return get_mmvq_mmid_max_batch_gcn(type);
#elif !defined(GGML_USE_MUSA) && (__CUDA_ARCH__ == GGML_CUDA_CC_VOLTA || __CUDA_ARCH__ >= GGML_CUDA_CC_ADA_LOVELACE)
#elif !defined(GGML_USE_MUSA) && (__CUDA_ARCH__ == GGML_CUDA_CC_PASCAL || __CUDA_ARCH__ == GGML_CUDA_CC_VOLTA || __CUDA_ARCH__ >= GGML_CUDA_CC_ADA_LOVELACE)
return MMVQ_MAX_BATCH_SIZE;
#elif !defined(GGML_USE_MUSA) && __CUDA_ARCH__ >= GGML_CUDA_CC_TURING
return get_mmvq_mmid_max_batch_turing_plus(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);
+29 -14
View File
@@ -149,7 +149,7 @@ __global__ void __launch_bounds__(d_state, 1)
const int src0_nb2, const int src0_nb3, const int src1_nb2, const int src1_nb3,
const int src2_nb1, const int src2_nb2, const int src3_nb1,
const int src4_nb2, const int src4_nb3, const int src5_nb2, const int src5_nb3,
const int64_t s_off, const int64_t n_head, const int64_t d_head, const int64_t n_group, const int64_t n_tok, const int64_t K) {
char * s_base, const int64_t s_slot_bytes, const int64_t n_head, const int64_t d_head, const int64_t n_group, const int64_t n_tok, const int64_t K) {
const float * GGML_CUDA_RESTRICT src0 = src0_ptr;
const float * GGML_CUDA_RESTRICT src1 = src1_ptr;
const float * GGML_CUDA_RESTRICT src2 = src2_ptr;
@@ -184,7 +184,7 @@ __global__ void __launch_bounds__(d_state, 1)
const float * B_warp = (const float *) ((const char *) src4 + (seq_idx * src4_nb3) + (group_off));
const float * C_warp = (const float *) ((const char *) src5 + (seq_idx * src5_nb3) + (group_off));
float * y_warp = dst + (seq_idx * n_tok * n_head * d_head) + warp_idx;
float * s_warp = (float *) ((char *) dst + s_off + seq_idx * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
float * s_warp = (float *) (s_base + seq_idx * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
// strides across n_seq_tokens
const int stride_x = src1_nb2 / sizeof(float);
@@ -227,7 +227,7 @@ __global__ void __launch_bounds__(d_state, 1)
// Slot 0 is the final state written below; slots 1..K-1 are rollback snapshots.
const int64_t slot = n_tok - 1 - i;
if (K > 1 && slot > 0 && slot < K) {
float * s_snapshot_warp = (float *) ((char *) dst + s_off + (slot * gridDim.y + seq_idx) * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
float * s_snapshot_warp = (float *) ((char *) s_warp + slot * s_slot_bytes);
#pragma unroll
for (int j = 0; j < c_factor; j++) {
s_snapshot_warp[WARP_SIZE * j + lane] = state[j];
@@ -248,7 +248,11 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
const int src2_nb2, const int src3_nb1, const int src4_nb2, const int src4_nb3, const int src5_nb2,
const int src5_nb3, const int64_t s_off, const int64_t d_state, const int64_t head_dim,
const int64_t n_head, const int64_t n_group, const int64_t n_tok, const int64_t n_seq,
const int64_t K, cudaStream_t stream) {
const int64_t K, const ggml_cuda_ssm_scan_fused_cache * cache, cudaStream_t stream) {
// when fused, the states go straight into the recurrent cache and the dst tail is left alone
char * const s_base = cache ? (char *) cache->data : (char *) dst + s_off;
const int64_t s_slot_bytes = cache ? cache->slot_stride * (int64_t) sizeof(float) : n_seq * (int64_t) src0_nb3;
// NOTE: if you change conditions here, be sure to update the corresponding supports_op condition!
if (src3_nb1 == sizeof(float)) {
// Mamba-2
@@ -261,7 +265,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
ggml_cuda_kernel_launch(ssm_scan_f32_group<96/WARP_SIZE, 96>, launch_params,
src0, src1, src2, src3, src4, src5, src6, dst,
src0_nb2, src0_nb3, src1_nb2, src1_nb3, src2_nb1, src2_nb2, src3_nb1,
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_base, s_slot_bytes, n_head, head_dim, n_group, n_tok, K);
} else if (d_state == 128) {
constexpr int threads = 128;
constexpr int num_warps = threads/WARP_SIZE;
@@ -271,7 +275,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
ggml_cuda_kernel_launch(ssm_scan_f32_group<128/WARP_SIZE, 128>, launch_params,
src0, src1, src2, src3, src4, src5, src6, dst,
src0_nb2, src0_nb3, src1_nb2, src1_nb3, src2_nb1, src2_nb2, src3_nb1,
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_base, s_slot_bytes, n_head, head_dim, n_group, n_tok, K);
} else if (d_state == 256) { // Falcon-H1
constexpr int threads = 256;
constexpr int num_warps = threads/WARP_SIZE;
@@ -281,7 +285,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
ggml_cuda_kernel_launch(ssm_scan_f32_group<256/WARP_SIZE, 256>, launch_params,
src0, src1, src2, src3, src4, src5, src6, dst,
src0_nb2, src0_nb3, src1_nb2, src1_nb3, src2_nb1, src2_nb2, src3_nb1,
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_base, s_slot_bytes, n_head, head_dim, n_group, n_tok, K);
} else {
GGML_ABORT("doesn't support d_state!=(96, 128 or 256).");
}
@@ -570,12 +574,13 @@ __global__ void ssm_ssd_scale_state_kernel(
}
// Copy initial state from src0[ids[s]] into s_cur for each sequence.
// src0 and s_cur can alias when the state is written straight into the cache.
// Grid: (ceil(d_state * head_dim * n_head / BLOCK), n_seqs)
template <int BLOCK_SIZE>
__global__ void ssm_ssd_init_state_kernel(
const float * __restrict__ src0, // {d_state, head_dim, n_head, n_rs}
const float * src0, // {d_state, head_dim, n_head, n_rs}
const int32_t * __restrict__ ids, // {n_seqs}
float * __restrict__ s_cur, // {d_state, head_dim, n_head, n_seqs}
float * s_cur, // {d_state, head_dim, n_head, n_seqs}
const int state_size, // d_state * head_dim * n_head
const int64_t s0_stride_seq) { // elements between state rows
const int s = blockIdx.y;
@@ -599,7 +604,8 @@ static void ssm_scan_ssd_f32_cuda(
const int A_stride, // A (src3) stride between heads
const int B_stride_tok, const int B_stride_seq, // B (src4) strides
const int C_stride_tok, const int C_stride_seq, // C (src5) strides
const int64_t s_off, const int64_t d_state, const int64_t head_dim,
float * s_cur, // state: dst state tail, or the cache when fused
const int64_t d_state, const int64_t head_dim,
const int64_t n_head, const int64_t n_group, const int64_t n_tok, const int64_t n_seq) {
cudaStream_t stream = ctx.stream();
@@ -625,7 +631,6 @@ static void ssm_scan_ssd_f32_cuda(
matmul_t * X_dt = X_dt_buf.get();
matmul_t * B_weighted = B_w_buf.get();
float * C_scaled = C_s_buf.get();
float * s_cur = (float *)((char *)dst_d + s_off); // write state directly to dst
// Step 1: softplus(dt) and parallel prefix sum over full sequence
{
@@ -780,7 +785,8 @@ static void ssm_scan_ssd_f32_cuda(
}
#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
static void ggml_cuda_op_ssm_scan_impl(ggml_backend_cuda_context & ctx, ggml_tensor * dst,
const ggml_cuda_ssm_scan_fused_cache * cache) {
const struct ggml_tensor * src0 = dst->src[0]; // s
const struct ggml_tensor * src1 = dst->src[1]; // x
const struct ggml_tensor * src2 = dst->src[2]; // dt
@@ -864,12 +870,21 @@ void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
(int)(src3->nb[1] / sizeof(float)),
(int)(src4->nb[2] / sizeof(float)), (int)(src4->nb[3] / sizeof(float)),
(int)(src5->nb[2] / sizeof(float)), (int)(src5->nb[3] / sizeof(float)),
s_off, nc, nr, nh, ng, n_t, n_s);
cache ? cache->data : (float *) ((char *) dst_d + s_off), nc, nr, nh, ng, n_t, n_s);
return;
}
#endif
ssm_scan_f32_cuda(src0_d, src1_d, src2_d, src3_d, src4_d, src5_d, src6_d, dst_d,
src0->nb[2], src0->nb[3], src1->nb[2], src1->nb[3], src2->nb[1], src2->nb[2],
src3->nb[1], src4->nb[2], src4->nb[3], src5->nb[2], src5->nb[3],
s_off, nc, nr, nh, ng, n_t, n_s, K, stream);
s_off, nc, nr, nh, ng, n_t, n_s, K, cache, stream);
}
void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
ggml_cuda_op_ssm_scan_impl(ctx, dst, nullptr);
}
void ggml_cuda_op_ssm_scan_fused_cache(ggml_backend_cuda_context & ctx, ggml_tensor * dst,
ggml_cuda_ssm_scan_fused_cache cache) {
ggml_cuda_op_ssm_scan_impl(ctx, dst, &cache);
}
+10
View File
@@ -1,3 +1,13 @@
#include "common.cuh"
// fused-kernel recurrent-state output; strides in elements (per-seq stride is always the state row size, set in-kernel)
struct ggml_cuda_ssm_scan_fused_cache {
float * data; // rollback slot 0
int64_t slot_stride; // between rollback slots
};
void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst);
// same op, but writes the state snapshot(s) into the cache instead of dst (see ggml_cuda_try_ssm_scan_cache_fusion)
void ggml_cuda_op_ssm_scan_fused_cache(ggml_backend_cuda_context & ctx, ggml_tensor * dst,
ggml_cuda_ssm_scan_fused_cache cache);
+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
}
+51 -8
View File
@@ -107,7 +107,7 @@ static __device__ __forceinline__ float op_ceil(float x) {
}
static __device__ __forceinline__ float op_round(float x) {
return round(x);
return roundf(x);
}
static __device__ __forceinline__ float op_trunc(float x) {
@@ -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);
}
}
}
+244 -23
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);
}
@@ -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;
}
}
}
+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; \
+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);
+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>;

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